nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
4.19 kB
"""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"]