File size: 3,652 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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
"""Public inference API and CLI for CPM-jev."""

from __future__ import annotations

import argparse
import json
from pathlib import Path
from typing import Any, Sequence

import torch
from transformers import AutoTokenizer

from model import BASE_MODEL_ID, MiniCPMJEVModel


def format_candidate(state: Any, kind: str, question: str, option: Any) -> str:
    if not isinstance(state, str):
        state = json.dumps(state, ensure_ascii=False, sort_keys=True)
    if not isinstance(option, str):
        option = json.dumps(option, ensure_ascii=False, sort_keys=True)
    return (
        f"State:\n{state}\n\n"
        f"Question type: {kind}\nQuestion: {question}\nCandidate option: {option}\n"
        "How well does this candidate answer the question?"
    )


class DecisionModel:
    """Scores candidate options and returns raw softmax probabilities."""

    def __init__(
        self,
        model_dir: str | Path,
        *,
        base_model: str = BASE_MODEL_ID,
        device: str | None = None,
        max_length: int = 512,
    ):
        self.model_dir = Path(model_dir)
        self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
        self.max_length = max_length
        self.tokenizer = AutoTokenizer.from_pretrained(self.model_dir, trust_remote_code=True)
        self.tokenizer.truncation_side = "left"
        if self.tokenizer.pad_token_id is None:
            self.tokenizer.pad_token = self.tokenizer.eos_token
        dtype = torch.bfloat16 if self.device.type == "cuda" else torch.float32
        self.model = MiniCPMJEVModel.from_pretrained(
            self.model_dir, base_model=base_model, dtype=dtype
        ).to(self.device).eval()

    @torch.inference_mode()
    def decide(
        self,
        *,
        state: Any,
        question: str,
        options: Sequence[Any],
        kind: str = "choice",
    ) -> dict[str, Any]:
        options = list(options)
        if len(options) < 2:
            raise ValueError("options must contain at least two candidates")
        texts = [format_candidate(state, kind, question, option) for option in options]
        encoded = self.tokenizer(
            texts,
            padding=True,
            truncation=True,
            max_length=self.max_length,
            return_tensors="pt",
        ).to(self.device)
        logits = self.model(encoded["input_ids"], encoded["attention_mask"]).float()
        probabilities = torch.softmax(logits, dim=-1).cpu().tolist()
        best = max(range(len(options)), key=probabilities.__getitem__)
        return {
            "options": options,
            "probabilities": probabilities,
            "choice": options[best],
            "confidence": probabilities[best],
        }


def main() -> None:
    parser = argparse.ArgumentParser(description="Run one CPM-jev decision")
    parser.add_argument("--model-dir", default=".")
    parser.add_argument("--base-model", default=BASE_MODEL_ID)
    parser.add_argument("--state", required=True)
    parser.add_argument("--question", required=True)
    parser.add_argument("--options", nargs="+", required=True)
    parser.add_argument("--kind", default="choice", choices=("choice", "noul", "score"))
    parser.add_argument("--device", default=None)
    args = parser.parse_args()
    model = DecisionModel(
        args.model_dir, base_model=args.base_model, device=args.device
    )
    result = model.decide(
        state=args.state,
        question=args.question,
        options=args.options,
        kind=args.kind,
    )
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()