nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
5.18 kB
"""Deterministic evaluation-time interventions for predicate memory.
These interventions are not training features. They are causal probes for asking
whether the learned predicate-memory path is example-specific, span-anchored, and
predicate-gated, rather than just extra residual capacity.
"""
from __future__ import annotations
from typing import Literal
import torch
PredicateMemoryIntervention = Literal[
"none",
"typed_graph_memory",
"zero_memory",
"memory_zero",
"uniform_gate",
"untyped_same_topology",
"span_roll",
"span_corrupted",
"span_reverse",
"batch_shuffle",
"batch_shuffled_graph",
"random_same_degree",
]
PREDICATE_MEMORY_INTERVENTIONS: tuple[str, ...] = (
"none",
"typed_graph_memory",
"zero_memory",
"memory_zero",
"uniform_gate",
"untyped_same_topology",
"span_roll",
"span_corrupted",
"span_reverse",
"batch_shuffle",
"batch_shuffled_graph",
"random_same_degree",
)
def _valid_lengths(attention_mask: torch.Tensor | None, batch: int, seq: int) -> list[int]:
if attention_mask is None:
return [seq] * batch
return [int(x) for x in attention_mask.to(torch.long).sum(dim=1).tolist()]
def _roll_valid(x: torch.Tensor, attention_mask: torch.Tensor | None, *, shifts: int = 1) -> torch.Tensor:
out = x.clone()
b, s = x.shape[:2]
for i, n in enumerate(_valid_lengths(attention_mask, b, s)):
if n > 1:
out[i, :n] = torch.roll(x[i, :n], shifts=shifts, dims=0)
return out
def _reverse_valid(x: torch.Tensor, attention_mask: torch.Tensor | None) -> torch.Tensor:
out = x.clone()
b, s = x.shape[:2]
for i, n in enumerate(_valid_lengths(attention_mask, b, s)):
if n > 1:
out[i, :n] = torch.flip(x[i, :n], dims=(0,))
return out
def _random_permute_valid(x: torch.Tensor, attention_mask: torch.Tensor | None) -> torch.Tensor:
"""Deterministically permute valid memory slots while preserving the budget."""
out = x.clone()
b, s = x.shape[:2]
for i, n in enumerate(_valid_lengths(attention_mask, b, s)):
if n > 1:
generator = torch.Generator(device="cpu")
generator.manual_seed(1729 + i * 1_000_003 + n * 97)
perm = torch.randperm(n, generator=generator).to(x.device)
out[i, :n] = x[i, perm]
return out
def apply_predicate_memory_intervention(
*,
candidate_key: torch.Tensor,
candidate_value: torch.Tensor,
predicate_gate: torch.Tensor,
attention_mask: torch.Tensor | None,
intervention: str = "none",
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Return intervened ``(key, value, gate)`` tensors.
``zero_memory`` makes the predicate-memory update exactly zero by zeroing the
values read by attention. The other interventions preserve tensor shapes and
most marginal statistics while corrupting a specific part of the learned
predicate-memory computation.
"""
if intervention in {"none", "typed_graph_memory"}:
return candidate_key, candidate_value, predicate_gate
if intervention not in PREDICATE_MEMORY_INTERVENTIONS:
raise ValueError(f"unknown predicate memory intervention {intervention!r}")
if intervention in {"zero_memory", "memory_zero"}:
return candidate_key, torch.zeros_like(candidate_value), predicate_gate
if intervention in {"uniform_gate", "untyped_same_topology"}:
gate = torch.ones_like(predicate_gate)
if attention_mask is not None:
gate = gate * attention_mask.to(gate.dtype).unsqueeze(-1)
return candidate_key, candidate_value, gate
if intervention in {"span_roll", "span_corrupted"}:
return (
_roll_valid(candidate_key, attention_mask),
_roll_valid(candidate_value, attention_mask),
_roll_valid(predicate_gate, attention_mask),
)
if intervention == "span_reverse":
return (
_reverse_valid(candidate_key, attention_mask),
_reverse_valid(candidate_value, attention_mask),
_reverse_valid(predicate_gate, attention_mask),
)
if intervention in {"batch_shuffle", "batch_shuffled_graph"}:
if candidate_key.shape[0] > 1:
return (
torch.roll(candidate_key, shifts=1, dims=0),
torch.roll(candidate_value, shifts=1, dims=0),
torch.roll(predicate_gate, shifts=1, dims=0),
)
return (
_roll_valid(candidate_key, attention_mask),
_roll_valid(candidate_value, attention_mask),
_roll_valid(predicate_gate, attention_mask),
)
if intervention == "random_same_degree":
return (
_random_permute_valid(candidate_key, attention_mask),
_random_permute_valid(candidate_value, attention_mask),
_random_permute_valid(predicate_gate, attention_mask),
)
raise AssertionError(f"unhandled intervention {intervention!r}")
__all__ = [
"PREDICATE_MEMORY_INTERVENTIONS",
"PredicateMemoryIntervention",
"apply_predicate_memory_intervention",
]