| """Exact execution of typed compositional programs over sparse graphs.""" |
|
|
| 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 |
| from strata.modeling.compose.ast import AnchorRef, Apply, ProgramNode |
| from strata.modeling.compose.types import GraphOperator, SemanticType, signature |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class TypedGraphValue: |
| value_type: SemanticType |
| node_state: torch.Tensor |
| filler: torch.Tensor | None = None |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class TypedAlgebraGraph: |
| relations: SparseRelationOperators |
| role_binder: OrthonormalRoleBinder |
| event_memories: torch.Tensor |
| node_values: torch.Tensor |
| temporal_rank: torch.Tensor |
|
|
| def __post_init__(self) -> None: |
| nodes = self.relations.node_count |
| if self.event_memories.shape != ( |
| nodes, |
| self.role_binder.filler_dim, |
| self.role_binder.role_dim, |
| ): |
| raise ValueError("event_memories have an incompatible shape") |
| if self.node_values.shape != (nodes, self.role_binder.filler_dim): |
| raise ValueError("node_values have an incompatible shape") |
| if self.temporal_rank.shape != (nodes,): |
| raise ValueError("temporal_rank must contain one value per node") |
|
|
|
|
| class TypedProgramExecutor: |
| """Execute a sound AST; learning never participates in graph traversal.""" |
|
|
| def execute( |
| self, |
| program: ProgramNode, |
| *, |
| anchors: dict[int, TypedGraphValue], |
| graph: TypedAlgebraGraph, |
| semiring: str = "max_product", |
| ) -> TypedGraphValue: |
| if isinstance(program, AnchorRef): |
| value = anchors[program.index] |
| if value.value_type is not program.value_type: |
| raise TypeError("runtime anchor type does not match its AST declaration") |
| return value |
|
|
| argument = self.execute(program.argument, anchors=anchors, graph=graph, semiring=semiring) |
| spec = signature(program.operator) |
| if argument.value_type is not spec.input_type: |
| raise TypeError("well-typed AST produced an incompatible runtime value") |
|
|
| if program.operator in (GraphOperator.LATEST_EVENT, GraphOperator.EARLIEST_EVENT): |
| state = self._select_temporal(argument.node_state, graph.temporal_rank, latest=( |
| program.operator is GraphOperator.LATEST_EVENT |
| )) |
| return TypedGraphValue(spec.output_type, state) |
|
|
| if spec.relation_name is None: |
| raise RuntimeError(f"operator {program.operator.value} has no execution rule") |
| state = graph.relations.step(argument.node_state, spec.relation_name, semiring=semiring) |
| filler = None |
| if spec.role_name is not None: |
| memory = torch.einsum("...n,nfd->...fd", argument.node_state, graph.event_memories) |
| filler = graph.role_binder.unbind(memory, spec.role_name) |
| |
| node_filler = torch.einsum("...n,nf->...f", state, graph.node_values) |
| if not torch.allclose(filler, node_filler, atol=1e-6, rtol=1e-6): |
| raise RuntimeError(f"{spec.role_name} relation and TPR memory disagree") |
| return TypedGraphValue(spec.output_type, state, filler) |
|
|
| @staticmethod |
| def _select_temporal(state: torch.Tensor, rank: torch.Tensor, *, latest: bool) -> torch.Tensor: |
| active = state > 0 |
| sentinel = -torch.inf if latest else torch.inf |
| scores = rank.to(device=state.device, dtype=state.dtype).expand_as(state) |
| scores = torch.where(active, scores, torch.full_like(scores, sentinel)) |
| index = scores.argmax(dim=-1) if latest else scores.argmin(dim=-1) |
| valid = active.any(dim=-1) |
| result = torch.zeros_like(state) |
| result.scatter_(-1, index.unsqueeze(-1), valid.to(state.dtype).unsqueeze(-1)) |
| return result |
|
|
|
|
| def permute_graph( |
| graph: TypedAlgebraGraph, |
| old_to_new: torch.Tensor, |
| ) -> TypedAlgebraGraph: |
| """Apply a graph isomorphism while preserving all relation semantics.""" |
|
|
| permutation = old_to_new.to(dtype=torch.long, device=graph.event_memories.device) |
| nodes = graph.relations.node_count |
| if permutation.shape != (nodes,) or sorted(permutation.tolist()) != list(range(nodes)): |
| raise ValueError("old_to_new must be a node permutation") |
| edges: dict[str, list[tuple[int, int, float]]] = { |
| name: [] for name in graph.relations.relation_names |
| } |
| for relation_id, source, target, weight in zip( |
| graph.relations.edge_relations.tolist(), |
| graph.relations.edge_sources.tolist(), |
| graph.relations.edge_targets.tolist(), |
| graph.relations.edge_weights.tolist(), |
| strict=True, |
| ): |
| edges[graph.relations.relation_names[relation_id]].append( |
| (int(permutation[source]), int(permutation[target]), float(weight)) |
| ) |
| relations = SparseRelationOperators(nodes, graph.relations.relation_names, edges).to(permutation.device) |
| event_memories = torch.empty_like(graph.event_memories) |
| node_values = torch.empty_like(graph.node_values) |
| temporal_rank = torch.empty_like(graph.temporal_rank) |
| event_memories[permutation] = graph.event_memories |
| node_values[permutation] = graph.node_values |
| temporal_rank[permutation] = graph.temporal_rank |
| return TypedAlgebraGraph( |
| relations=relations, |
| role_binder=graph.role_binder, |
| event_memories=event_memories, |
| node_values=node_values, |
| temporal_rank=temporal_rank, |
| ) |
|
|
|
|
| def permute_value(value: TypedGraphValue, old_to_new: torch.Tensor) -> TypedGraphValue: |
| state = torch.empty_like(value.node_state) |
| state[..., old_to_new] = value.node_state |
| return TypedGraphValue(value.value_type, state, value.filler) |
|
|
|
|
| __all__ = [ |
| "TypedAlgebraGraph", |
| "TypedGraphValue", |
| "TypedProgramExecutor", |
| "permute_graph", |
| "permute_value", |
| ] |
|
|