"""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", ]