| """Factorized natural predicate, operator, and argument-span compiler.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
|
|
| import torch |
| from torch import nn |
|
|
|
|
| @dataclass(slots=True) |
| class NaturalCompilerOutput: |
| predicate_logits: torch.Tensor |
| operator_logits: torch.Tensor |
| span_start_logits: torch.Tensor |
| span_end_logits: torch.Tensor |
| argument_head_logits: torch.Tensor |
| role_presence_logits: torch.Tensor |
|
|
|
|
| class NaturalAtomicCompiler(nn.Module): |
| """Compile atomic natural queries without an answer or graph-state bypass.""" |
|
|
| def __init__( |
| self, |
| pretrained_word_vectors: torch.Tensor, |
| *, |
| hidden_dim: int = 192, |
| query_dim: int = 128, |
| role_count: int = 5, |
| ) -> None: |
| super().__init__() |
| if pretrained_word_vectors.ndim != 2: |
| raise ValueError("pretrained_word_vectors must be [vocabulary, width]") |
| self.role_count = int(role_count) |
| self.word_vectors = nn.Embedding.from_pretrained( |
| pretrained_word_vectors.to(torch.float32), freeze=True, padding_idx=0 |
| ) |
| self.word_projection = nn.Sequential( |
| nn.LayerNorm(pretrained_word_vectors.shape[1]), |
| nn.Linear(pretrained_word_vectors.shape[1], hidden_dim), |
| nn.GELU(), |
| ) |
| self.anchor_projection = nn.Sequential( |
| nn.LayerNorm(pretrained_word_vectors.shape[1]), |
| nn.Linear(pretrained_word_vectors.shape[1], hidden_dim), |
| nn.GELU(), |
| ) |
| self.anchor_scale = nn.Parameter(torch.tensor(10.0)) |
| self.context_encoder = nn.GRU( |
| hidden_dim, |
| hidden_dim, |
| num_layers=2, |
| batch_first=True, |
| bidirectional=True, |
| dropout=0.1, |
| ) |
| context_dim = 2 * hidden_dim |
| self.predicate_head = nn.Linear(context_dim, 1) |
|
|
| self.query_projection = nn.Sequential( |
| nn.LayerNorm(pretrained_word_vectors.shape[1]), |
| nn.Linear(pretrained_word_vectors.shape[1], query_dim), |
| nn.GELU(), |
| ) |
| self.query_encoder = nn.GRU( |
| query_dim, |
| query_dim, |
| batch_first=True, |
| bidirectional=True, |
| ) |
| self.operator_head = nn.Sequential( |
| nn.LayerNorm(2 * query_dim), |
| nn.Linear(2 * query_dim, context_dim), |
| nn.GELU(), |
| nn.Linear(context_dim, role_count), |
| ) |
|
|
| self.role_embedding = nn.Embedding(role_count, hidden_dim) |
| self.word_span_projection = nn.Linear(context_dim, hidden_dim, bias=False) |
| self.predicate_span_projection = nn.Linear(context_dim, hidden_dim, bias=False) |
| self.start_head = nn.Linear(hidden_dim, 1, bias=False) |
| self.end_head = nn.Linear(hidden_dim, 1, bias=False) |
| self.argument_head = nn.Linear(hidden_dim, 1, bias=False) |
| self.role_presence = nn.Sequential( |
| nn.Linear(context_dim + hidden_dim, hidden_dim), |
| nn.GELU(), |
| nn.Linear(hidden_dim, 1), |
| ) |
|
|
| def forward( |
| self, |
| word_ids: torch.Tensor, |
| word_mask: torch.Tensor, |
| query_word_ids: torch.Tensor, |
| query_mask: torch.Tensor, |
| *, |
| anchor_word_ids: torch.Tensor | None = None, |
| gold_predicate_index: torch.Tensor | None = None, |
| ) -> NaturalCompilerOutput: |
| vectors = self.word_projection(self.word_vectors(word_ids)) |
| lengths = word_mask.sum(dim=-1).to(torch.long).cpu() |
| packed = nn.utils.rnn.pack_padded_sequence( |
| vectors, lengths, batch_first=True, enforce_sorted=False |
| ) |
| encoded, _ = self.context_encoder(packed) |
| words, _ = nn.utils.rnn.pad_packed_sequence( |
| encoded, batch_first=True, total_length=word_ids.shape[1] |
| ) |
| predicate_logits = self.predicate_head(words).squeeze(-1) |
| if anchor_word_ids is not None: |
| anchor = self.anchor_projection(self.word_vectors(anchor_word_ids)) |
| lexical = torch.nn.functional.normalize(vectors, dim=-1) |
| anchor = torch.nn.functional.normalize(anchor, dim=-1) |
| predicate_logits = predicate_logits + self.anchor_scale.clamp(0.0, 20.0) * torch.einsum( |
| "bsd,bd->bs", lexical, anchor |
| ) |
| predicate_logits = predicate_logits.masked_fill(~word_mask, -torch.inf) |
| if gold_predicate_index is None: |
| predicate_weights = torch.softmax(predicate_logits, dim=-1) |
| predicate = torch.einsum("bs,bsd->bd", predicate_weights, words) |
| else: |
| predicate = words.gather( |
| 1, |
| gold_predicate_index[:, None, None].expand(-1, 1, words.shape[-1]), |
| ).squeeze(1) |
|
|
| query_lengths = query_mask.sum(dim=-1).to(torch.long).cpu() |
| query_packed = nn.utils.rnn.pack_padded_sequence( |
| self.query_projection(self.word_vectors(query_word_ids)), |
| query_lengths, |
| batch_first=True, |
| enforce_sorted=False, |
| ) |
| _query_words, query_hidden = self.query_encoder(query_packed) |
| query = torch.cat([query_hidden[-2], query_hidden[-1]], dim=-1) |
| operator_logits = self.operator_head(query) |
|
|
| role_ids = torch.arange(self.role_count, device=word_ids.device) |
| role = self.role_embedding(role_ids) |
| span_hidden = torch.tanh( |
| self.word_span_projection(words)[:, None, :, :] |
| + self.predicate_span_projection(predicate)[:, None, None, :] |
| + role[None, :, None, :] |
| ) |
| start = self.start_head(span_hidden).squeeze(-1).masked_fill(~word_mask[:, None, :], -torch.inf) |
| end = self.end_head(span_hidden).squeeze(-1).masked_fill(~word_mask[:, None, :], -torch.inf) |
| argument_head = self.argument_head(span_hidden).squeeze(-1).masked_fill( |
| ~word_mask[:, None, :], -torch.inf |
| ) |
| presence_input = torch.cat( |
| [ |
| predicate[:, None, :].expand(-1, self.role_count, -1), |
| role[None, :, :].expand(word_ids.shape[0], -1, -1), |
| ], |
| dim=-1, |
| ) |
| presence = self.role_presence(presence_input).squeeze(-1) |
| return NaturalCompilerOutput( |
| predicate_logits, operator_logits, start, end, argument_head, presence |
| ) |
|
|
|
|
| def decode_valid_spans( |
| start_logits: torch.Tensor, |
| end_logits: torch.Tensor, |
| word_mask: torch.Tensor, |
| *, |
| maximum_width: int = 24, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| """Return the highest scoring valid start/end pair for every role.""" |
|
|
| scores = start_logits.unsqueeze(-1) + end_logits.unsqueeze(-2) |
| length = start_logits.shape[-1] |
| positions = torch.arange(length, device=scores.device) |
| valid = positions[:, None] <= positions[None, :] |
| valid &= positions[None, :] - positions[:, None] < maximum_width |
| valid = valid[None, None, :, :] & ( |
| word_mask[:, None, :, None] & word_mask[:, None, None, :] |
| ) |
| flat = scores.masked_fill(~valid, -torch.inf).flatten(-2) |
| selected = flat.argmax(dim=-1) |
| return selected // length, selected % length |
|
|
|
|
| def gated_compiler_reliability( |
| confidence: torch.Tensor, |
| *, |
| threshold: float, |
| ) -> torch.Tensor: |
| """Pass qualified compiler confidence to BRR and zero every unsafe read.""" |
|
|
| if not 0.0 <= threshold <= 1.0: |
| raise ValueError("threshold must be in [0, 1]") |
| finite = torch.isfinite(confidence) |
| qualified = finite & (confidence >= threshold) |
| return torch.where(qualified, confidence.clamp(0.0, 1.0), torch.zeros_like(confidence)) |
|
|
|
|
| __all__ = [ |
| "NaturalAtomicCompiler", |
| "NaturalCompilerOutput", |
| "decode_valid_spans", |
| "gated_compiler_reliability", |
| ] |
|
|