strata-headquotient-q25 / source /src /strata /modeling /ph_pat /predicate_hypergraph.py
nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
6.6 kB
"""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"]