| """Compile the predicted three-stream state consumed by the LM path.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
|
|
| import torch |
| from torch import nn |
|
|
| from strata.modeling.ph_pat.config import PHPATArm, PHPATConfig |
| from strata.modeling.ph_pat.dependency_chart import DependencyChart, DependencyChartOutput |
| from strata.modeling.ph_pat.predicate_hypergraph import PredicateHypergraphCompiler, PredicateHypergraphState |
| from strata.modeling.ph_pat.primitive_registers import PrimitiveRegisterState |
| from strata.modeling.ph_pat.role_primitive_bridge import BridgeOutput, RolePrimitiveBridge |
| from strata.modeling.ph_pat.segment_commit import SegmentLayout |
| from strata.modeling.ph_pat.span_compiler import SpanCompiler, SpanCompilerOutput |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class CompiledPHPATChart: |
| spans: SpanCompilerOutput |
| dependency: DependencyChartOutput |
| hypergraph: PredicateHypergraphState |
| primitives: PrimitiveRegisterState | None |
| bridge_permutation: torch.Tensor |
| arm: str |
|
|
|
|
| class ChartExecutor(nn.Module): |
| def __init__(self, config: PHPATConfig) -> None: |
| super().__init__() |
| self.config = config |
| self.span_compiler = SpanCompiler(config) |
| self.dependency_chart = DependencyChart(config) |
| self.hypergraph = PredicateHypergraphCompiler(config) |
| self.bridge = RolePrimitiveBridge(config) |
|
|
| def forward( |
| self, |
| hidden: torch.Tensor, |
| layout: SegmentLayout, |
| *, |
| arm: PHPATArm | None = None, |
| corruption: str = "none", |
| generator: torch.Generator | None = None, |
| ) -> CompiledPHPATChart: |
| selected = arm or self.config.arm |
| spans = self.span_compiler(hidden, layout) |
| dependency = self.dependency_chart(spans) |
| typed_event = selected in {PHPATArm.TYPED_EVENT, PHPATArm.FULL, PHPATArm.SHUFFLED_BRIDGE, PHPATArm.RANDOM_CHART} |
| typed_primitive = selected in {PHPATArm.TYPED_PRIMITIVE, PHPATArm.FULL, PHPATArm.SHUFFLED_BRIDGE, PHPATArm.RANDOM_CHART} |
| hypergraph = self.hypergraph(hidden, layout, dependency, typed_roles=typed_event) |
|
|
| if selected == PHPATArm.GENERIC_MEMORY: |
| pooled = hypergraph.frames.mean(dim=2, keepdim=True) |
| hypergraph = hypergraph.with_frames(pooled.expand_as(hypergraph.frames)) |
| elif selected == PHPATArm.RANDOM_CHART or corruption == "random_matched": |
| frames = _online_random_matched(hypergraph.frames, generator=generator) |
| hypergraph = hypergraph.with_frames(frames) |
|
|
| use_primitives = selected in { |
| PHPATArm.TYPED_PRIMITIVE, |
| PHPATArm.FULL, |
| PHPATArm.RANDOM_CHART, |
| PHPATArm.SHUFFLED_BRIDGE, |
| } |
| bridge: BridgeOutput | None = None |
| if use_primitives: |
| bridge = self.bridge( |
| hypergraph, |
| typed_primitives=typed_primitive, |
| shuffle_bridge=selected == PHPATArm.SHUFFLED_BRIDGE or corruption == "bridge_shuffle", |
| generator=generator, |
| ) |
| permutation = ( |
| bridge.role_permutation |
| if bridge is not None |
| else torch.arange(self.config.role_count, device=hidden.device) |
| ) |
| return CompiledPHPATChart( |
| spans=spans, |
| dependency=dependency, |
| hypergraph=hypergraph, |
| primitives=bridge.primitives if bridge is not None else None, |
| bridge_permutation=permutation, |
| arm=selected.value, |
| ) |
|
|
|
|
| def _online_random_matched(tensor: torch.Tensor, *, generator: torch.Generator | None) -> torch.Tensor: |
| sample = torch.randn(tensor.shape, device=tensor.device, dtype=torch.float32, generator=generator) |
| source = tensor.float() |
| reduce_dims = tuple(range(1, tensor.ndim)) |
| mean = source.mean(dim=reduce_dims, keepdim=True) |
| std = source.std(dim=reduce_dims, keepdim=True).clamp_min(1e-6) |
| sample = (sample - sample.mean(dim=reduce_dims, keepdim=True)) / sample.std( |
| dim=reduce_dims, keepdim=True |
| ).clamp_min(1e-6) |
| return (sample * std + mean).to(tensor.dtype) |
|
|
|
|
| __all__ = ["ChartExecutor", "CompiledPHPATChart"] |
|
|