nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
6.14 kB
"""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",
]