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