nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
12.7 kB
"""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",
]