File size: 3,017 Bytes
03223d7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
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()