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