Text Classification
jev-style
Safetensors
minicpm
minicpm5
system-one
decision-model
probability
calibration
agent
routing
Instructions to use link921/CPM-jev with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- jev-style
How to use link921/CPM-jev with jev-style:
pip install "jev-style[torch]"
from jev_style import JevStyle, noul, choice js = JevStyle.from_pretrained("link921/CPM-jev") out = js.decide("I was charged twice for one order.", { "billing": noul("This message is about billing."), "team": choice("Which team should handle it?", ["billing", "shipping", "tech"]), }) print(out["answers"]["team"]["choice"]) - Notebooks
- Google Colab
- Kaggle
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()
|