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 [""])[: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