nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
5.95 kB
"""Exact tensor-product role binding with a fixed orthonormal role basis."""
from __future__ import annotations
from collections.abc import Sequence
import torch
from torch import nn
DEFAULT_ROLE_NAMES = (
"ARG0",
"ARG1",
"ARG2",
"TIME",
"LOCATION",
"OTHER",
"NULL",
)
class OrthonormalRoleBinder(nn.Module):
"""Bind filler vectors to canonical role vectors and unbind them exactly.
The memory matrix is ``sum_r filler[r] outer role[r]``. Canonical one-hot
role vectors avoid numerical drift and make the Phase 0 contract exact.
A supplied orthonormal basis can be used for controlled permutation tests.
"""
def __init__(
self,
filler_dim: int,
role_names: Sequence[str] = DEFAULT_ROLE_NAMES,
*,
role_dim: int | None = None,
basis: torch.Tensor | None = None,
) -> None:
super().__init__()
if filler_dim <= 0:
raise ValueError("filler_dim must be positive")
names = tuple(str(name) for name in role_names)
if not names or len(set(names)) != len(names):
raise ValueError("role_names must be non-empty and unique")
role_dim = len(names) if role_dim is None else int(role_dim)
if role_dim < len(names):
raise ValueError("role_dim must be at least the number of roles")
if basis is None:
basis = torch.eye(role_dim, dtype=torch.float32)[: len(names)]
else:
basis = basis.detach().to(dtype=torch.float32)
if basis.shape != (len(names), role_dim):
raise ValueError(
f"basis must have shape {(len(names), role_dim)}, got {tuple(basis.shape)}"
)
gram = basis @ basis.transpose(0, 1)
identity = torch.eye(len(names), dtype=gram.dtype, device=gram.device)
if not torch.allclose(gram, identity, atol=1e-6, rtol=1e-6):
raise ValueError("role basis must be orthonormal")
self.filler_dim = int(filler_dim)
self.role_names = names
self.role_dim = role_dim
self.register_buffer("basis", basis.contiguous(), persistent=True)
@property
def role_count(self) -> int:
return len(self.role_names)
def role_id(self, name: str) -> int:
try:
return self.role_names.index(name)
except ValueError as exc:
raise KeyError(f"unknown role {name!r}") from exc
def role_vector(self, roles: int | str | torch.Tensor) -> torch.Tensor:
if isinstance(roles, str):
roles = self.role_id(roles)
if isinstance(roles, int):
if not 0 <= roles < self.role_count:
raise IndexError("role index out of range")
return self.basis[roles]
role_ids = roles.to(device=self.basis.device, dtype=torch.long)
if bool(((role_ids < 0) | (role_ids >= self.role_count)).any()):
raise IndexError("role index out of range")
return self.basis[role_ids]
def bind(self, fillers: torch.Tensor, *, role_mask: torch.Tensor | None = None) -> torch.Tensor:
"""Return a TPR matrix from ``[..., roles, filler_dim]`` fillers."""
if fillers.shape[-2:] != (self.role_count, self.filler_dim):
raise ValueError(
"fillers must end with "
f"[{self.role_count}, {self.filler_dim}], got {tuple(fillers.shape)}"
)
values = fillers
if role_mask is not None:
expected = fillers.shape[:-1]
if role_mask.shape != expected:
raise ValueError(f"role_mask must have shape {expected}, got {tuple(role_mask.shape)}")
values = values * role_mask.to(device=values.device, dtype=values.dtype).unsqueeze(-1)
basis = self.basis.to(device=values.device, dtype=values.dtype)
return torch.einsum("...rf,rd->...fd", values, basis)
def unbind(self, memory: torch.Tensor, roles: int | str | torch.Tensor) -> torch.Tensor:
"""Apply the requested role as an algebraic unbinding operator."""
if memory.shape[-2:] != (self.filler_dim, self.role_dim):
raise ValueError(
"memory must end with "
f"[{self.filler_dim}, {self.role_dim}], got {tuple(memory.shape)}"
)
role_vector = self.role_vector(roles).to(device=memory.device, dtype=memory.dtype)
if role_vector.ndim == 1:
return torch.einsum("...fd,d->...f", memory, role_vector)
if role_vector.shape[:-1] != memory.shape[:-2]:
raise ValueError(
"tensor role IDs must match memory batch dimensions: "
f"{tuple(role_vector.shape[:-1])} != {tuple(memory.shape[:-2])}"
)
return torch.einsum("...fd,...d->...f", memory, role_vector)
def unbind_all(self, memory: torch.Tensor) -> torch.Tensor:
if memory.shape[-2:] != (self.filler_dim, self.role_dim):
raise ValueError("memory has incompatible trailing dimensions")
basis = self.basis.to(device=memory.device, dtype=memory.dtype)
return torch.einsum("...fd,rd->...rf", memory, basis)
def permuted(self, permutation: torch.Tensor) -> "OrthonormalRoleBinder":
"""Return a binder with globally relabeled role basis vectors."""
permutation = permutation.to(dtype=torch.long, device=self.basis.device)
if permutation.shape != (self.role_count,):
raise ValueError(f"permutation must have shape ({self.role_count},)")
if sorted(permutation.tolist()) != list(range(self.role_count)):
raise ValueError("permutation must contain each role exactly once")
return OrthonormalRoleBinder(
self.filler_dim,
self.role_names,
role_dim=self.role_dim,
basis=self.basis[permutation].cpu(),
)
__all__ = ["DEFAULT_ROLE_NAMES", "OrthonormalRoleBinder"]