File size: 4,190 Bytes
7c5e40e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
"""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"]