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