| """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 |
| node_mask: torch.Tensor |
| edge_type: torch.Tensor |
| edge_mask: torch.Tensor |
| edge_src: torch.Tensor |
| edge_dst: torch.Tensor |
| edge_rel: torch.Tensor |
| edge_slot_mask: torch.Tensor |
|
|
|
|
| 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", |
| ] |
|
|