gump2049's picture
Initial private MLX 8bit release
084eaf7 verified
Raw History Blame Contribute Delete
3.39 kB
"""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))