"""Sparse query-program compilation without free evidence-slot selection.""" from __future__ import annotations from dataclasses import dataclass from enum import IntEnum import torch from torch import nn from strata.modeling.algebra.entmax import entmax_bisect, sparse_topk class GraphOperation(IntEnum): NO_GRAPH_READ = 0 READ_ENTITY = 1 READ_EVENT_ARG0 = 2 READ_EVENT_ARG1 = 3 READ_EVENT_TIME = 4 FOLLOW_COREF = 5 FOLLOW_TEMPORAL_BEFORE = 6 FOLLOW_SUPPORT = 7 FOLLOW_CONTRADICTION = 8 class AnchorType(IntEnum): ENTITY = 0 EVENT = 1 SEGMENT = 2 TEMPORAL = 3 @dataclass(frozen=True, slots=True) class QueryProgramOutput: query_state: torch.Tensor anchor_type_logits: torch.Tensor anchor_scores: torch.Tensor operation_logits: torch.Tensor operation_probabilities: torch.Tensor reliability: torch.Tensor @property def hard_anchor(self) -> torch.Tensor: return self.anchor_scores.argmax(dim=-1) @property def hard_operations(self) -> torch.Tensor: return self.operation_logits.argmax(dim=-1) class QueryProgramCompiler(nn.Module): """Compile frozen token states into an anchor type and short operation list. Candidate scoring is set-equivariant: candidate order is never embedded. The scorer predicts graph anchors only; evidence slots do not enter this module. Operation probabilities use alpha-entmax and are capped at two active operations per step. """ def __init__( self, d_model: int, *, hidden_dim: int = 256, max_steps: int = 3, maximum_active_operations: int = 2, entmax_alpha: float = 1.5, ) -> None: super().__init__() if d_model <= 0 or hidden_dim <= 0 or max_steps <= 0: raise ValueError("dimensions and max_steps must be positive") if not 1 <= maximum_active_operations <= len(GraphOperation): raise ValueError("invalid maximum_active_operations") self.d_model = int(d_model) self.max_steps = int(max_steps) self.maximum_active_operations = int(maximum_active_operations) self.entmax_alpha = float(entmax_alpha) self.query_attention = nn.Linear(d_model, 1, bias=False) self.query_projection = nn.Linear(d_model, d_model, bias=False) self.candidate_projection = nn.Linear(d_model, d_model, bias=False) self.anchor_scorer = nn.Sequential( nn.Linear(4 * d_model, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, 1), ) self.anchor_type_head = nn.Linear(d_model, len(AnchorType)) self.operation_head = nn.Sequential( nn.Linear(d_model, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, max_steps * len(GraphOperation)), ) def forward( self, hidden_states: torch.Tensor, *, query_mask: torch.Tensor, candidate_states: torch.Tensor, candidate_mask: torch.Tensor, ) -> QueryProgramOutput: if hidden_states.ndim != 3 or hidden_states.shape[-1] != self.d_model: raise ValueError("hidden_states must have shape [batch, sequence, d_model]") if query_mask.shape != hidden_states.shape[:2]: raise ValueError("query_mask must have shape [batch, sequence]") if candidate_states.ndim != 3 or candidate_states.shape[0] != hidden_states.shape[0]: raise ValueError("candidate_states must have shape [batch, candidates, d_model]") if candidate_states.shape[-1] != self.d_model or candidate_mask.shape != candidate_states.shape[:2]: raise ValueError("candidate state width or mask is incompatible") attention_logits = self.query_attention(hidden_states).squeeze(-1) attention_logits = attention_logits.masked_fill(~query_mask.bool(), -torch.inf) attention = torch.softmax(attention_logits.float(), dim=-1).to(hidden_states.dtype) attention = torch.where(query_mask.any(dim=-1, keepdim=True), attention, torch.zeros_like(attention)) query = torch.einsum("bs,bsd->bd", attention, hidden_states) q = self.query_projection(query) candidates = self.candidate_projection(candidate_states) expanded = q.unsqueeze(1).expand_as(candidates) anchor_features = torch.cat( [expanded, candidates, expanded * candidates, candidates - expanded], dim=-1 ) anchor_scores = self.anchor_scorer(anchor_features).squeeze(-1) anchor_scores = anchor_scores.masked_fill(~candidate_mask.bool(), -torch.inf) operation_logits = self.operation_head(query).view( hidden_states.shape[0], self.max_steps, len(GraphOperation) ) operation_probabilities = entmax_bisect( operation_logits.float(), alpha=self.entmax_alpha, dim=-1 ).to(operation_logits.dtype) operation_probabilities = sparse_topk( operation_probabilities, k=self.maximum_active_operations, dim=-1, ) anchor_confidence = torch.softmax(anchor_scores.float(), dim=-1).amax(dim=-1) operation_confidence = operation_probabilities.float().amax(dim=-1).mean(dim=-1) reliability = (anchor_confidence * operation_confidence).clamp(0, 1).detach() return QueryProgramOutput( query_state=query, anchor_type_logits=self.anchor_type_head(query), anchor_scores=anchor_scores, operation_logits=operation_logits, operation_probabilities=operation_probabilities, reliability=reliability, ) __all__ = [ "AnchorType", "GraphOperation", "QueryProgramCompiler", "QueryProgramOutput", ]