File size: 2,460 Bytes
257d034
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)