File size: 3,392 Bytes
084eaf7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Run APUS-OpenJev MLX on Apple Silicon: exact candidate distribution with mlx-lm.

    pip install mlx-lm
    python examples/openjev_mlx.py --model apus-ailab/APUS-OpenJev-v1-4B-MLX-8bit

The prompt is rendered by ``openjev_contracts.py`` (identical to the training
contract) inside the no-thinking chat turn; the candidate labels A..P are single
tokens and the distribution is the softmax of their final-position logits.
"""

import argparse
import json
import sys
from pathlib import Path

import mlx.core as mx
from mlx_lm import load
from mlx_lm.models.cache import make_prompt_cache

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from openjev_contracts import format_response, label_mapping, render_prompt  # noqa: E402

MAX_TOKENS = 8192
PREFILL_STEP = 512  # chunked prefill keeps long prompts within unified memory


class OpenJevMLX:
    def __init__(self, model_path):
        self.model, self.tokenizer = load(model_path)

    def compile(self, request):
        mapping = label_mapping(request)
        prompt = self.tokenizer.apply_chat_template(
            [{"role": "user", "content": render_prompt(request)}],
            tokenize=False,
            add_generation_prompt=True,
            enable_thinking=False,
        )
        ids = self.tokenizer.encode(prompt, add_special_tokens=False)
        if not ids or len(ids) > MAX_TOKENS:
            raise ValueError("input exceeds 8192 tokens; no truncation is performed")
        candidates = []
        for label in mapping:
            token = self.tokenizer.encode(label, add_special_tokens=False)
            if (
                len(token) != 1
                or self.tokenizer.encode(prompt + label, add_special_tokens=False)
                != ids + token
            ):
                raise ValueError(
                    "candidate label is not a single token at the answer boundary"
                )
            candidates.append(token[0])
        return ids, candidates

    def decide(self, request):
        ids, candidates = self.compile(request)
        cache = make_prompt_cache(self.model)
        for start in range(0, len(ids) - 1, PREFILL_STEP):
            self.model(
                mx.array([ids[start : min(start + PREFILL_STEP, len(ids) - 1)]]),
                cache=cache,
            )
            mx.eval([c.state for c in cache])
        logits = self.model(mx.array([ids[-1:]]), cache=cache)[0, -1].astype(
            mx.float32
        )[mx.array(candidates)]
        response = format_response(request, mx.softmax(logits).tolist())
        response.update(prompt_tokens=len(ids), calibrated=False)
        return response


EXAMPLE = {
    "id": "support-731",
    "group_id": "support-731",
    "primitive": "choice",
    "state": "Order 731 was delivered. The customer confirms that the issue is resolved.",
    "instructions": "Select the next support action.",
    "criteria": [
        {"id": "close_ticket", "description": "Close the ticket as resolved."},
        {"id": "escalate", "description": "Escalate to a human agent."},
        {"id": "refund", "description": "Issue a refund."},
    ],
}

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", required=True, help="local directory or HF repo id")
    args = parser.parse_args()
    print(json.dumps(OpenJevMLX(args.model).decide(EXAMPLE), indent=2))