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