openjev-e4b / openjev /decision_head.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
3.02 kB
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()