"""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 and TPR paths are required to agree for deterministic one-hot graphs. 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", ]