"""Short typed graph programs with exact role unbinding.""" from __future__ import annotations from dataclasses import dataclass import torch from strata.modeling.algebra.relations import SparseRelationOperators from strata.modeling.algebra.role_binding import OrthonormalRoleBinder @dataclass(frozen=True, slots=True) class GraphProgram: relations: tuple[str, ...] = () read_role: str | None = None def __post_init__(self) -> None: if len(self.relations) > 3: raise ValueError("qualification graph programs are bounded to three hops") @dataclass(frozen=True, slots=True) class GraphProgramResult: node_state: torch.Tensor filler: torch.Tensor | None trace: tuple[torch.Tensor, ...] class GraphProgramExecutor: def __init__( self, relations: SparseRelationOperators, role_binder: OrthonormalRoleBinder, ) -> None: self.relations = relations self.role_binder = role_binder def execute( self, anchor_state: torch.Tensor, program: GraphProgram, *, event_memories: torch.Tensor | None = None, semiring: str = "max_product", ) -> GraphProgramResult: state = anchor_state trace = [state] for relation in program.relations: state = self.relations.step(state, relation, semiring=semiring) trace.append(state) filler = None if program.read_role is not None: if event_memories is None: raise ValueError("event_memories are required for a role read") expected = ( *state.shape[:-1], self.relations.node_count, self.role_binder.filler_dim, self.role_binder.role_dim, ) if event_memories.shape != expected: raise ValueError(f"event_memories must have shape {expected}, got {tuple(event_memories.shape)}") memory = torch.einsum("...n,...nfd->...fd", state, event_memories) filler = self.role_binder.unbind(memory, program.read_role) return GraphProgramResult(node_state=state, filler=filler, trace=tuple(trace)) __all__ = ["GraphProgram", "GraphProgramExecutor", "GraphProgramResult"]