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