nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
3.08 kB
"""Anchor-linking interfaces kept separate from graph-program compilation."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from torch import nn
from strata.modeling.compose.ast import AnchorRef
from strata.modeling.compose.types import SemanticType
from strata.modeling.participant_late_interaction import QueryFacetCrossAttentionReader
@dataclass(frozen=True, slots=True)
class AnchorLinkOutput:
scores: torch.Tensor
selected_index: torch.Tensor
anchor_type: SemanticType
class StructuredAnchorLinker(nn.Module):
"""Copy an externally verified episode-local anchor without classifying IDs."""
def forward(self, local_indices: torch.Tensor, *, anchor_type: SemanticType) -> AnchorLinkOutput:
if local_indices.ndim != 1:
raise ValueError("local_indices must have shape [batch]")
candidates = int(local_indices.max().item()) + 1
scores = torch.full(
(local_indices.shape[0], candidates),
-torch.inf,
device=local_indices.device,
)
scores.scatter_(1, local_indices.unsqueeze(1), 0.0)
return AnchorLinkOutput(scores, local_indices, anchor_type)
class ParticipantFacetAnchorLinker(nn.Module):
"""Adapt the frozen four-facet participant primitive to entity anchors."""
def __init__(self, reader: QueryFacetCrossAttentionReader) -> None:
super().__init__()
self.reader = reader
def forward(
self,
query_tokens: torch.Tensor,
query_mask: torch.Tensor,
query_vector: torch.Tensor,
participant_facets: torch.Tensor,
) -> AnchorLinkOutput:
scores = self.reader(query_tokens, query_mask, query_vector, participant_facets)
return AnchorLinkOutput(scores, scores.argmax(dim=-1), SemanticType.ENTITY)
class TypedCandidateAnchorLinker(nn.Module):
"""Set-equivariant event/claim linker; candidate IDs are never embedded."""
def __init__(self, hidden_dim: int, anchor_type: SemanticType) -> None:
super().__init__()
self.anchor_type = anchor_type
self.query = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.candidate = nn.Linear(hidden_dim, hidden_dim, bias=False)
def forward(
self,
query: torch.Tensor,
candidates: torch.Tensor,
candidate_mask: torch.Tensor | None = None,
) -> AnchorLinkOutput:
scores = torch.einsum("bd,bnd->bn", self.query(query), self.candidate(candidates))
scores = scores / query.shape[-1] ** 0.5
if candidate_mask is not None:
scores = scores.masked_fill(~candidate_mask, -torch.inf)
return AnchorLinkOutput(scores, scores.argmax(dim=-1), self.anchor_type)
def copied_anchor(output: AnchorLinkOutput, batch_index: int = 0) -> AnchorRef:
return AnchorRef(int(output.selected_index[batch_index]), output.anchor_type)
__all__ = [
"AnchorLinkOutput",
"ParticipantFacetAnchorLinker",
"StructuredAnchorLinker",
"TypedCandidateAnchorLinker",
"copied_anchor",
]