Pivot / pivot_model.py
Q1z's picture
Pivot H200 full retrain
14bf8c2 verified
Raw History Blame Contribute Delete
8.02 kB
from __future__ import annotations
import hashlib
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
def masked_mean_pool(hidden:torch.Tensor, mask:torch.Tensor)->torch.Tensor:
m=mask.unsqueeze(-1).to(hidden.dtype)
return (hidden*m).sum(1)/m.sum(1).clamp_min(1e-6)
def _fingerprint_tensor(t:torch.Tensor, n:int=4096)->str:
x=t.detach().view(-1)[:n].float().cpu().contiguous().numpy().tobytes()
return hashlib.sha256(x).hexdigest()
class HashTokenizer:
is_hash_tokenizer=True
pad_token_id=0
def __init__(self,vocab_size=256): self.vocab_size=vocab_size
def __call__(self,texts,*,padding,max_length,truncation,return_tensors):
rows=[]; masks=[]
for text in texts:
toks=(str(text).lower().split() or ["<empty>"])[:max_length]
ids=[2+int(hashlib.md5(x.encode()).hexdigest()[:8],16)%(self.vocab_size-2) for x in toks]
mask=[1]*len(ids); ids += [0]*(max_length-len(ids)); mask += [0]*(max_length-len(mask))
rows.append(ids); masks.append(mask)
return {"input_ids":torch.tensor(rows),"attention_mask":torch.tensor(masks)}
class StubEncoder(nn.Module):
def __init__(self,vocab_size=256,hidden_size=32):
super().__init__()
self.hidden_size=hidden_size
self.embed=nn.Embedding(vocab_size,hidden_size,padding_idx=0)
self.pretrained_audit={"stub":True}
def forward(self,input_ids,attention_mask):
return masked_mean_pool(self.embed(input_ids),attention_mask)
class HuggingFaceEncoder(nn.Module):
"""Load the official MLM wrapper and extract its pretrained bidirectional LFM2 body."""
def __init__(self, model_id:str, revision:str, trust_remote_code:bool=True):
super().__init__()
from transformers import AutoModelForMaskedLM
loaded, info = AutoModelForMaskedLM.from_pretrained(
model_id,
revision=revision,
trust_remote_code=trust_remote_code,
torch_dtype=torch.float32,
output_loading_info=True,
low_cpu_mem_usage=False,
attn_implementation='sdpa',
)
if not hasattr(loaded,"lfm2"):
raise RuntimeError(f"Expected masked-LM wrapper with .lfm2, got {type(loaded).__name__}")
missing=[str(x) for x in info.get("missing_keys",[])]
unexpected=[str(x) for x in info.get("unexpected_keys",[])]
fatal_missing=[x for x in missing if x.startswith("lfm2.")]
fatal_unexpected=[x for x in unexpected if x.startswith("lfm2.")]
if fatal_missing or fatal_unexpected:
raise RuntimeError(f"Pretrained backbone load failed: missing={fatal_missing[:20]} unexpected={fatal_unexpected[:20]}")
self.model=loaded.lfm2
self.hidden_size=int(self.model.config.hidden_size)
if "Bidirectional" not in type(self.model).__name__:
raise RuntimeError(f"Expected bidirectional encoder, got {type(self.model).__name__}")
causal=[name for name,m in self.model.named_modules() if hasattr(m,"is_causal") and bool(getattr(m,"is_causal"))]
if causal:
raise RuntimeError(f"Causal attention survived in encoder: {causal[:20]}")
named=dict(self.model.named_parameters())
probes=[k for k in ["embed_tokens.weight","layers.0.feed_forward.w1.weight","layers.0.self_attn.q_proj.weight"] if k in named]
if not probes:
probes=list(named)[:3]
parameter_count=sum(p.numel() for p in self.model.parameters())
if parameter_count < 300_000_000:
raise RuntimeError(f"Unexpectedly small pretrained encoder: {parameter_count} parameters")
self.pretrained_audit={
"wrapper_class":type(loaded).__name__,
"encoder_class":type(self.model).__name__,
"missing_keys":missing,
"unexpected_keys":unexpected,
"parameter_count":parameter_count,
"fingerprints":{k:_fingerprint_tensor(named[k]) for k in probes},
}
del loaded
def forward(self,input_ids,attention_mask):
out=self.model(input_ids=input_ids,attention_mask=attention_mask,use_cache=False)
h=out.last_hidden_state if hasattr(out,"last_hidden_state") else out[0]
return masked_mean_pool(h,attention_mask)
class MLPScorer(nn.Module):
"""Pivot scorer: preserve the original DSBT matching contract."""
def __init__(self,hidden_size,mlp_hidden):
super().__init__()
self.net=nn.Sequential(
nn.Linear(3*hidden_size,mlp_hidden),
nn.GELU(),
nn.Linear(mlp_hidden,1),
)
def forward(self,hc,ho):
hc2=hc.unsqueeze(1).expand_as(ho)
return self.net(torch.cat([hc2,ho,hc2*ho],-1)).squeeze(-1)
@dataclass
class SetBrierOutput:
logits:torch.Tensor
probs:torch.Tensor
pred_index:torch.Tensor
h_c:torch.Tensor
h_o:torch.Tensor
class SetBrierEncoder(nn.Module):
def __init__(self,encoder,scorer):
super().__init__(); self.encoder=encoder; self.scorer=scorer
def encoder_parameters(self): return self.encoder.parameters()
def scorer_parameters(self): return self.scorer.parameters()
def encode(self,ids,mask): return self.encoder(ids,mask)
def score_preencoded(self,hc,ho,opt_mask=None):
if hc.ndim==1: hc=hc.unsqueeze(0)
if ho.ndim==2: ho=ho.unsqueeze(0)
if hc.ndim!=2 or ho.ndim!=3: raise ValueError("bad embedding shapes")
if opt_mask is None:
opt_mask=torch.ones(ho.shape[:2],dtype=torch.bool,device=ho.device)
ho_real=ho*opt_mask.unsqueeze(-1).to(ho.dtype)
logits=self.scorer(hc,ho_real)
logits=logits.masked_fill(~opt_mask.bool(),torch.finfo(logits.dtype).min)
probs=F.softmax(logits.float(),dim=-1)
return SetBrierOutput(logits,probs,probs.argmax(-1),hc,ho)
def forward(self,ctx_ids,ctx_mask,opt_ids,opt_mask,opt_attn):
hc=self.encode(ctx_ids,ctx_mask)
b,k,l=opt_ids.shape
flat_ids=opt_ids.reshape(b*k,l)
flat_attn=opt_attn.reshape(b*k,l)
real=opt_mask.reshape(-1).bool()
if not bool(real.any().item()):
raise ValueError("batch has no real options")
# Pad option slots are mathematically excluded from the scorer/softmax.
# Do not waste H200 compute encoding them.
real_h=self.encode(flat_ids[real],flat_attn[real])
flat_h=real_h.new_zeros((b*k,real_h.shape[-1]))
flat_h=flat_h.index_copy(0,real.nonzero(as_tuple=False).squeeze(1),real_h)
ho=flat_h.reshape(b,k,-1)
return self.score_preencoded(hc,ho,opt_mask)
@torch.no_grad()
def decide(self,ctx_ids,ctx_mask,opt_ids,opt_mask,opt_attn,option_texts):
self.eval()
out=self.forward(ctx_ids,ctx_mask,opt_ids,opt_mask,opt_attn)
probs=out.probs[0][opt_mask[0].bool()].detach().cpu().float()
probs=probs/probs.sum().clamp_min(1e-12)
index=int(out.pred_index[0].item())
return {"choice":option_texts[index],"index":index,"probs":probs.tolist()}
def build_tokenizer(cfg):
bb=cfg["backbone"]
if bb.get("encoder")=="stub":
return HashTokenizer(int(bb.get("stub_vocab_size",256)))
from transformers import AutoTokenizer
return AutoTokenizer.from_pretrained(bb["model_id"],revision=bb["revision"],trust_remote_code=True)
def build_model(cfg,device=None):
device=device or torch.device("cpu")
bb=cfg["backbone"]
if bb.get("encoder")=="stub":
enc=StubEncoder(int(bb.get("stub_vocab_size",256)),int(bb.get("stub_hidden_size",32)))
else:
enc=HuggingFaceEncoder(bb["model_id"],bb["revision"],bool(bb.get("trust_remote_code",True)))
scorer=MLPScorer(enc.hidden_size,int(cfg["scorer"].get("hidden_size",enc.hidden_size)))
model=SetBrierEncoder(enc,scorer).to(device)
if next(model.parameters()).dtype != torch.float32:
raise RuntimeError("Expected FP32 master parameters")
return model