File size: 6,597 Bytes
e69b72a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""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)
        # EVENT is an address, not an independently pooled generic message.
        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"]