File size: 6,597 Bytes
e69b72a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | """Predicted event anchors and hard typed event-role memory addresses."""
from __future__ import annotations
from dataclasses import dataclass
import math
import torch
from torch import nn
from strata.modeling.ph_pat.config import PHPATConfig
from strata.modeling.ph_pat.dependency_chart import DependencyChartOutput
from strata.modeling.ph_pat.segment_commit import SegmentLayout
@dataclass(frozen=True, slots=True)
class PredicateHypergraphState:
event_nodes: torch.Tensor
frames: torch.Tensor
event_segment_ids: torch.Tensor
event_valid: torch.Tensor
trigger_logits: torch.Tensor
role_filler_logits: torch.Tensor
role_filler_weights: torch.Tensor
def read(self, event_ids: torch.Tensor, role_ids: torch.Tensor) -> torch.Tensor:
"""Hard hierarchical read: frames[event_id][role_id]."""
if event_ids.shape != role_ids.shape:
raise ValueError("event_ids and role_ids must have identical shapes")
batch = self.frames.shape[0]
if event_ids.shape[0] != batch:
raise ValueError("address batch does not match frame batch")
if bool((event_ids < 0).any() or (event_ids >= self.frames.shape[1]).any()):
raise IndexError("event address out of range")
if bool((role_ids < 0).any() or (role_ids >= self.frames.shape[2]).any()):
raise IndexError("role address out of range")
batch_index = torch.arange(batch, device=self.frames.device)
view = (batch,) + (1,) * (event_ids.ndim - 1)
return self.frames[batch_index.view(view), event_ids, role_ids]
def with_frames(self, frames: torch.Tensor) -> "PredicateHypergraphState":
if frames.shape != self.frames.shape:
raise ValueError("replacement frames must preserve shape")
return PredicateHypergraphState(
event_nodes=self.event_nodes,
frames=frames,
event_segment_ids=self.event_segment_ids,
event_valid=self.event_valid,
trigger_logits=self.trigger_logits,
role_filler_logits=self.role_filler_logits,
role_filler_weights=self.role_filler_weights,
)
class PredicateHypergraphCompiler(nn.Module):
def __init__(self, config: PHPATConfig) -> None:
super().__init__()
self.config = config
self.event_queries = nn.Parameter(torch.empty(config.events_per_segment, config.d_model))
self.event_q = nn.Linear(config.d_model, config.d_model, bias=False)
self.token_k = nn.Linear(config.d_model, config.d_model, bias=False)
self.token_v = nn.Linear(config.d_model, config.d_model, bias=False)
self.trigger = nn.Linear(config.d_model, config.events_per_segment)
self.role_embeddings = nn.Parameter(torch.empty(config.role_count, config.d_model))
self.role_query = nn.Linear(config.d_model, config.d_model, bias=False)
self.filler_key = nn.Linear(config.d_model, config.d_model, bias=False)
self.role_values = nn.ModuleList(
nn.Linear(config.d_model, config.d_model, bias=False) for _ in range(config.role_count)
)
nn.init.normal_(self.event_queries, std=0.02)
nn.init.normal_(self.role_embeddings, std=0.02)
def forward(
self,
hidden: torch.Tensor,
layout: SegmentLayout,
chart: DependencyChartOutput,
*,
typed_roles: bool,
) -> PredicateHypergraphState:
batch, _seq, dim = hidden.shape
event_blocks: list[torch.Tensor] = []
valid_blocks: list[torch.Tensor] = []
segment_blocks: list[torch.Tensor] = []
token_keys = self.token_k(hidden)
token_values = self.token_v(hidden)
for segment in range(layout.segment_count):
queries = self.event_q(self.event_queries).view(1, self.config.events_per_segment, dim).expand(batch, -1, -1)
scores = torch.matmul(queries, token_keys.transpose(-2, -1)) / math.sqrt(dim)
mask = layout.segment_token_mask[:, segment].unsqueeze(1)
scores = scores.masked_fill(~mask, torch.finfo(scores.dtype).min)
weights = torch.softmax(scores.float(), dim=-1).to(hidden.dtype)
weights = torch.nan_to_num(weights)
event_blocks.append(torch.matmul(weights, token_values))
valid_blocks.append(layout.segment_valid[:, segment].unsqueeze(1).expand(batch, self.config.events_per_segment))
segment_blocks.append(torch.full((batch, self.config.events_per_segment), segment, device=hidden.device, dtype=torch.long))
events = torch.cat(event_blocks, dim=1)
event_valid = torch.cat(valid_blocks, dim=1)
event_segments = torch.cat(segment_blocks, dim=1)
role_embeddings = self.role_embeddings
if not typed_roles:
role_embeddings = role_embeddings.mean(dim=0, keepdim=True).expand_as(role_embeddings)
role_query = self.role_query(events.unsqueeze(2) + role_embeddings.view(1, 1, self.config.role_count, dim))
filler_key = self.filler_key(chart.nodes)
logits = torch.einsum("berd,bnd->bern", role_query, filler_key) / math.sqrt(dim)
same_segment = event_segments.unsqueeze(-1) == chart.segment_ids.unsqueeze(1)
valid = event_valid.unsqueeze(-1) & chart.valid_mask.unsqueeze(1) & same_segment
logits = logits.masked_fill(~valid.unsqueeze(2), torch.finfo(logits.dtype).min)
weights = torch.softmax(logits.float(), dim=-1).to(hidden.dtype)
weights = torch.nan_to_num(weights)
shared_value = sum(layer(chart.nodes) for layer in self.role_values) / self.config.role_count
frame_values: list[torch.Tensor] = []
for role, projection in enumerate(self.role_values):
values = projection(chart.nodes) if typed_roles else shared_value
frame_values.append(torch.einsum("ben,bnd->bed", weights[:, :, role], values))
frames = torch.stack(frame_values, dim=2)
# EVENT is an address, not an independently pooled generic message.
frames = frames.clone()
frames[:, :, 0] = events
frames = frames * event_valid.unsqueeze(-1).unsqueeze(-1).to(frames.dtype)
return PredicateHypergraphState(
event_nodes=events,
frames=frames,
event_segment_ids=event_segments,
event_valid=event_valid,
trigger_logits=self.trigger(hidden),
role_filler_logits=logits,
role_filler_weights=weights,
)
__all__ = ["PredicateHypergraphCompiler", "PredicateHypergraphState"]
|