"""MiniCPM5 backbone with a scalar option-scoring head for CPM-jev.""" from __future__ import annotations from pathlib import Path from typing import Any import torch from peft import PeftModel from safetensors.torch import load_file from torch import nn from transformers import AutoModel BASE_MODEL_ID = "openbmb/MiniCPM5-2B-Base" class MiniCPMJEVModel(nn.Module): """A MiniCPM5 encoder backbone plus a learned scalar candidate scorer.""" def __init__(self, backbone: nn.Module, hidden_size: int): super().__init__() self.backbone = backbone self.decision_head = nn.Linear(hidden_size, 1, bias=True) @staticmethod def _hidden_size(config: Any) -> int: text_config = config.get_text_config() if hasattr(config, "get_text_config") else config for name in ("hidden_size", "dim", "n_embd"): if hasattr(text_config, name): return int(getattr(text_config, name)) raise ValueError("Cannot determine the MiniCPM hidden size from its config") @classmethod def from_pretrained( cls, model_dir: str | Path, *, base_model: str = BASE_MODEL_ID, dtype: torch.dtype | None = None, ) -> "MiniCPMJEVModel": model_dir = Path(model_dir) if dtype is None: dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 base = AutoModel.from_pretrained(base_model, dtype=dtype, trust_remote_code=True) backbone = PeftModel.from_pretrained(base, str(model_dir), is_trainable=False) model = cls(backbone, cls._hidden_size(base.config)) head_path = model_dir / "decision_head.safetensors" model.decision_head.load_state_dict(load_file(str(head_path))) return model def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=False, return_dict=True, ) hidden = outputs.last_hidden_state positions = torch.arange(hidden.shape[1], device=hidden.device).unsqueeze(0) last = positions.masked_fill(attention_mask.eq(0), -1).max(dim=1).values.clamp_min(0) pooled = hidden[torch.arange(hidden.shape[0], device=hidden.device), last] return self.decision_head(pooled.to(self.decision_head.weight.dtype)).squeeze(-1)