Zero-Shot Classification
Safetensors
PEFT
English
openjev
classification
decision-model
listwise
gemma4
research
Instructions to use bambamdevs/openjev-e4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use bambamdevs/openjev-e4b with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
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()
|