"""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()