File size: 12,728 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
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
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
"""First-class token-anchored graph-object utilities.

The v1 graph object is intentionally minimal and dense:

* ``node_type`` / ``node_mask`` over token positions.
* ``edge_type`` / ``edge_mask`` over ``[query_token, key_token]`` pairs.

This mirrors the current UD/SRL supervision contract and gives the model an
explicit graph source that can later be compared under typed, untyped,
relation-permuted, random, span-corrupted, and zero-object modes at matched
tensor budget.
"""

from __future__ import annotations

from typing import Literal, TypedDict

import torch

from strata.data.srl_labels import NONE_LOCAL, SRL_BASE
from strata.data.ud_labels import IGNORE_INDEX

IGNORE_GRAPH_LABEL = -100

GraphObjectIntervention = Literal[
    "none",
    "typed_graph_object",
    "untyped_same_topology",
    "relation_permuted",
    "random_same_degree",
    "span_corrupted",
    "batch_shuffled_graph",
    "zero_graph_object",
]

GRAPH_OBJECT_INTERVENTIONS: tuple[str, ...] = (
    "none",
    "typed_graph_object",
    "untyped_same_topology",
    "relation_permuted",
    "random_same_degree",
    "span_corrupted",
    "batch_shuffled_graph",
    "zero_graph_object",
)


class GraphObject(TypedDict):
    node_type: torch.Tensor  # [B, S], long, IGNORE_GRAPH_LABEL when absent
    node_mask: torch.Tensor  # [B, S], bool
    edge_type: torch.Tensor  # [B, S, S], long, IGNORE_GRAPH_LABEL when absent
    edge_mask: torch.Tensor  # [B, S, S], bool
    edge_src: torch.Tensor  # [B, E], long token index
    edge_dst: torch.Tensor  # [B, E], long token index
    edge_rel: torch.Tensor  # [B, E], long relation id
    edge_slot_mask: torch.Tensor  # [B, E], bool


def graph_object_from_batch(batch: dict[str, torch.Tensor]) -> GraphObject | None:
    """Extract a graph object from a collated batch if present."""

    required = (
        "graph_node_type",
        "graph_node_mask",
        "graph_edge_type",
        "graph_edge_mask",
        "graph_edge_src",
        "graph_edge_dst",
        "graph_edge_rel",
        "graph_edge_slot_mask",
    )
    if not all(key in batch for key in required):
        return None
    return {
        "node_type": batch["graph_node_type"],
        "node_mask": batch["graph_node_mask"],
        "edge_type": batch["graph_edge_type"],
        "edge_mask": batch["graph_edge_mask"],
        "edge_src": batch["graph_edge_src"],
        "edge_dst": batch["graph_edge_dst"],
        "edge_rel": batch["graph_edge_rel"],
        "edge_slot_mask": batch["graph_edge_slot_mask"],
    }


def graph_object_from_targets(batch: dict[str, torch.Tensor]) -> GraphObject | None:
    """Build a graph object from collated UD/SRL supervision tensors.

    This is used by eval loaders that do not go through ``collate_graph`` but do
    have gold arc/role targets. It preserves duplicate relations through edge
    slots while also filling a dense compatibility matrix.
    """

    if "input_ids" not in batch:
        return None
    input_ids = batch["input_ids"]
    b, s = input_ids.shape
    device = input_ids.device
    node_type = torch.full((b, s), IGNORE_GRAPH_LABEL, dtype=torch.long, device=device)
    node_mask = torch.zeros((b, s), dtype=torch.bool, device=device)
    dense_edge_type = torch.full((b, s, s), IGNORE_GRAPH_LABEL, dtype=torch.long, device=device)
    dense_edge_mask = torch.zeros((b, s, s), dtype=torch.bool, device=device)
    edge_lists: list[list[tuple[int, int, int]]] = [[] for _ in range(b)]

    if "node_type" in batch:
        node_type = batch["node_type"].clone()
        node_mask |= node_type != IGNORE_INDEX
    if "candidate_head" in batch:
        node_mask |= batch["candidate_head"].to(torch.bool)

    if "arc_head" in batch and "rel" in batch:
        arc_head = batch["arc_head"]
        rel = batch["rel"]
        for bi in range(b):
            supervised = ((arc_head[bi] != IGNORE_INDEX) & (rel[bi] != IGNORE_INDEX)).nonzero(as_tuple=False).flatten()
            for dep_t in supervised.tolist():
                head_t = int(arc_head[bi, dep_t].item())
                rel_t = int(rel[bi, dep_t].item())
                if 0 <= head_t < s:
                    edge_lists[bi].append((dep_t, head_t, rel_t))
                    dense_edge_type[bi, dep_t, head_t] = rel_t
                    dense_edge_mask[bi, dep_t, head_t] = True
                    node_mask[bi, dep_t] = True
                    node_mask[bi, head_t] = True

    if "predicate_mask" in batch:
        node_mask |= batch["predicate_mask"].to(torch.bool)
    if "role_target" in batch:
        role_target = batch["role_target"]
        for bi in range(b):
            active = ((role_target[bi] != IGNORE_INDEX) & (role_target[bi] != NONE_LOCAL)).nonzero(as_tuple=False)
            for arg_t, pred_t in active.tolist():
                rel_t = SRL_BASE + int(role_target[bi, arg_t, pred_t].item())
                edge_lists[bi].append((arg_t, pred_t, rel_t))
                dense_edge_type[bi, arg_t, pred_t] = rel_t
                dense_edge_mask[bi, arg_t, pred_t] = True
                node_mask[bi, arg_t] = True
                node_mask[bi, pred_t] = True

    max_edges = max(1, max((len(edges) for edges in edge_lists), default=0))
    edge_src = torch.zeros((b, max_edges), dtype=torch.long, device=device)
    edge_dst = torch.zeros((b, max_edges), dtype=torch.long, device=device)
    edge_rel = torch.full((b, max_edges), IGNORE_GRAPH_LABEL, dtype=torch.long, device=device)
    edge_slot_mask = torch.zeros((b, max_edges), dtype=torch.bool, device=device)
    for bi, edges in enumerate(edge_lists):
        for ei, (src, dst, rel_t) in enumerate(edges):
            edge_src[bi, ei] = src
            edge_dst[bi, ei] = dst
            edge_rel[bi, ei] = rel_t
            edge_slot_mask[bi, ei] = True
    if not bool(edge_slot_mask.any()):
        return None
    return {
        "node_type": node_type,
        "node_mask": node_mask,
        "edge_type": dense_edge_type,
        "edge_mask": dense_edge_mask,
        "edge_src": edge_src,
        "edge_dst": edge_dst,
        "edge_rel": edge_rel,
        "edge_slot_mask": edge_slot_mask,
    }


def _clone_graph_object(graph_object: GraphObject) -> GraphObject:
    return {
        "node_type": graph_object["node_type"].clone(),
        "node_mask": graph_object["node_mask"].clone(),
        "edge_type": graph_object["edge_type"].clone(),
        "edge_mask": graph_object["edge_mask"].clone(),
        "edge_src": graph_object["edge_src"].clone(),
        "edge_dst": graph_object["edge_dst"].clone(),
        "edge_rel": graph_object["edge_rel"].clone(),
        "edge_slot_mask": graph_object["edge_slot_mask"].clone(),
    }


def _valid_lengths(graph_object: GraphObject) -> list[int]:
    """Infer valid token prefix lengths from node/edge occupancy."""

    node_mask = graph_object["node_mask"].to(torch.bool)
    edge_mask = graph_object["edge_mask"].to(torch.bool)
    lengths: list[int] = []
    for i in range(node_mask.shape[0]):
        occupied = node_mask[i].clone()
        slots = graph_object["edge_slot_mask"][i].to(torch.bool)
        if bool(slots.any()):
            edge_src = graph_object["edge_src"][i, slots]
            edge_dst = graph_object["edge_dst"][i, slots]
            occupied[edge_src.clamp(0, occupied.numel() - 1)] = True
            occupied[edge_dst.clamp(0, occupied.numel() - 1)] = True
        if bool(edge_mask[i].any()):
            occupied |= edge_mask[i].any(dim=0)
            occupied |= edge_mask[i].any(dim=1)
        nz = occupied.nonzero(as_tuple=False).flatten()
        lengths.append(int(nz.max().item()) + 1 if nz.numel() else 0)
    return lengths


def _permute_graph_nodes(graph_object: GraphObject, permutations: list[torch.Tensor]) -> GraphObject:
    out = _clone_graph_object(graph_object)
    for batch_index, perm in enumerate(permutations):
        if perm.numel() <= 1:
            continue
        device = out["node_type"].device
        perm = perm.to(device)
        inv = torch.empty_like(perm)
        inv[perm] = torch.arange(perm.numel(), device=device)
        out["node_type"][batch_index, : perm.numel()] = graph_object["node_type"][batch_index, perm]
        out["node_mask"][batch_index, : perm.numel()] = graph_object["node_mask"][batch_index, perm]
        sub_edge_type = graph_object["edge_type"][batch_index, : perm.numel(), : perm.numel()]
        sub_edge_mask = graph_object["edge_mask"][batch_index, : perm.numel(), : perm.numel()]
        out["edge_type"][batch_index, : perm.numel(), : perm.numel()] = sub_edge_type[perm][:, perm]
        out["edge_mask"][batch_index, : perm.numel(), : perm.numel()] = sub_edge_mask[perm][:, perm]
        slots = graph_object["edge_slot_mask"][batch_index].to(torch.bool)
        if bool(slots.any()):
            src = graph_object["edge_src"][batch_index, slots].clamp(0, perm.numel() - 1)
            dst = graph_object["edge_dst"][batch_index, slots].clamp(0, perm.numel() - 1)
            out["edge_src"][batch_index, slots] = inv[src]
            out["edge_dst"][batch_index, slots] = inv[dst]
    return out


def apply_graph_object_intervention(
    graph_object: GraphObject | None,
    *,
    intervention: str = "none",
    relation_vocab_size: int,
) -> GraphObject | None:
    """Return an intervened graph object without changing tensor shapes."""

    if graph_object is None:
        return None
    if intervention in {"none", "typed_graph_object"}:
        return graph_object
    if intervention not in GRAPH_OBJECT_INTERVENTIONS:
        raise ValueError(f"unknown graph object intervention {intervention!r}")

    out = _clone_graph_object(graph_object)
    edge_mask = out["edge_mask"]

    if intervention == "zero_graph_object":
        out["node_mask"].zero_()
        out["edge_mask"].zero_()
        out["edge_slot_mask"].zero_()
        out["node_type"].fill_(IGNORE_GRAPH_LABEL)
        out["edge_type"].fill_(IGNORE_GRAPH_LABEL)
        out["edge_rel"].fill_(IGNORE_GRAPH_LABEL)
        return out

    if intervention == "untyped_same_topology":
        out["edge_type"] = out["edge_type"].masked_fill(edge_mask, 0)
        out["edge_rel"] = out["edge_rel"].masked_fill(out["edge_slot_mask"], 0)
        return out

    if intervention == "relation_permuted":
        typed = out["edge_type"].clamp_min(0)
        out["edge_type"] = out["edge_type"].masked_scatter(
            edge_mask,
            ((typed[edge_mask] + 1) % relation_vocab_size).to(out["edge_type"].dtype),
        )
        slot_mask = out["edge_slot_mask"]
        typed_slots = out["edge_rel"].clamp_min(0)
        out["edge_rel"] = out["edge_rel"].masked_scatter(
            slot_mask,
            ((typed_slots[slot_mask] + 1) % relation_vocab_size).to(out["edge_rel"].dtype),
        )
        return out

    if intervention == "batch_shuffled_graph":
        if out["node_type"].shape[0] > 1:
            return {
                "node_type": torch.roll(out["node_type"], shifts=1, dims=0),
                "node_mask": torch.roll(out["node_mask"], shifts=1, dims=0),
                "edge_type": torch.roll(out["edge_type"], shifts=1, dims=0),
                "edge_mask": torch.roll(out["edge_mask"], shifts=1, dims=0),
                "edge_src": torch.roll(out["edge_src"], shifts=1, dims=0),
                "edge_dst": torch.roll(out["edge_dst"], shifts=1, dims=0),
                "edge_rel": torch.roll(out["edge_rel"], shifts=1, dims=0),
                "edge_slot_mask": torch.roll(out["edge_slot_mask"], shifts=1, dims=0),
            }
        intervention = "span_corrupted"

    if intervention == "span_corrupted":
        lengths = _valid_lengths(out)
        permutations = [
            torch.roll(torch.arange(max(1, n), device=out["node_type"].device), shifts=1)
            for n in lengths
        ]
        return _permute_graph_nodes(graph_object, permutations)

    if intervention == "random_same_degree":
        permutations: list[torch.Tensor] = []
        for i, n in enumerate(_valid_lengths(out)):
            if n <= 1:
                permutations.append(torch.arange(max(1, n), device=out["node_type"].device))
                continue
            generator = torch.Generator(device="cpu")
            generator.manual_seed(92821 + i * 1_000_003 + n * 131)
            permutations.append(torch.randperm(n, generator=generator))
        return _permute_graph_nodes(graph_object, permutations)

    raise AssertionError(f"unhandled graph object intervention {intervention!r}")


__all__ = [
    "GRAPH_OBJECT_INTERVENTIONS",
    "IGNORE_GRAPH_LABEL",
    "GraphObject",
    "GraphObjectIntervention",
    "apply_graph_object_intervention",
    "graph_object_from_batch",
    "graph_object_from_targets",
]