"""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"]