from __future__ import annotations import math import torch import torch.nn.functional as F from torch import nn class PointerDecisionHead(nn.Module): def __init__(self, hidden_size: int, head_dim: int = 512, normalize: bool = True): super().__init__() self.normalize = normalize self.q_proj = nn.Linear(hidden_size, head_dim, bias=False) self.k_proj = nn.Linear(hidden_size, head_dim, bias=False) self.scale = math.sqrt(head_dim) def forward(self, decide_h: torch.Tensor, option_h: torch.Tensor, option_mask: torch.Tensor) -> torch.Tensor: decide_h = decide_h.float() option_h = option_h.float() q = self.q_proj(decide_h) k = self.k_proj(option_h) if self.normalize: q = F.normalize(q, dim=-1) k = F.normalize(k, dim=-1) logits = torch.einsum("bd,bnd->bn", q, k) * 10.0 else: logits = torch.einsum("bd,bnd->bn", q, k) / self.scale return logits.masked_fill(~option_mask, -1e9) class SpanDecisionHead(nn.Module): """Scores the mean hidden state across each option's semantic token span.""" def __init__(self, hidden_size: int, head_dim: int = 384): super().__init__() self.q_proj = nn.Linear(hidden_size, head_dim, bias=False) self.k_proj = nn.Linear(hidden_size, head_dim, bias=False) self.mlp = nn.Sequential( nn.Linear(head_dim * 4, head_dim), nn.GELU(), nn.Linear(head_dim, 1), ) def forward(self, decide_h: torch.Tensor, option_mean_h: torch.Tensor, option_mask: torch.Tensor) -> torch.Tensor: q = F.normalize(self.q_proj(decide_h.float()), dim=-1) k = F.normalize(self.k_proj(option_mean_h.float()), dim=-1) qe = q[:, None, :].expand_as(k) x = torch.cat([qe, k, qe * k, (qe - k).abs()], dim=-1) logits = self.mlp(x).squeeze(-1) return logits.masked_fill(~option_mask, -1e9) class HybridDecisionHead(nn.Module): """Pointer baseline plus a learned residual semantic-span scorer.""" def __init__(self, hidden_size: int, head_dim: int = 384, pointer_dim: int = 384, normalize: bool = True): super().__init__() self.pointer = PointerDecisionHead(hidden_size, pointer_dim, normalize=normalize) self.span = SpanDecisionHead(hidden_size, head_dim) self.logit_scale = nn.Parameter(torch.tensor(0.0)) def forward( self, decide_h: torch.Tensor, option_last_h: torch.Tensor, option_mean_h: torch.Tensor, option_mask: torch.Tensor, ) -> torch.Tensor: base = self.pointer(decide_h, option_last_h, option_mask) residual = self.span(decide_h, option_mean_h, option_mask) return base + self.logit_scale.exp().clamp(max=4.0) * residual def brier_loss(probs: torch.Tensor, targets: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: diff = (probs - targets).pow(2) * mask.float() return diff.sum(dim=-1).mean()