CPM-jev / model.py
link921's picture
Release v0.1 research preview files
257d034 verified
Raw History Blame Contribute Delete
2.46 kB
"""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)