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