"""Validate the existing full checkpoint on an Apple Silicon Mac with >=48 GiB.""" import argparse import importlib.util import importlib.metadata import json import logging import os from pathlib import Path import platform import subprocess import sys import threading import time import urllib.request def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--snapshot", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--target-prompt-tokens", type=int, default=86000) args = parser.parse_args() if platform.system() != "Darwin" or platform.machine() != "arm64": raise SystemExit("This validation requires a native Apple Silicon Mac.") physical_bytes = int(subprocess.check_output(["sysctl", "-n", "hw.memsize"])) if physical_bytes < 48 * 1024**3: raise SystemExit("Refusing a full checkpoint load: at least 48 GiB is required for this test.") if args.target_prompt_tokens < 4096: raise SystemExit("Use at least 4096 tokens for the full-model follow-up test.") args.snapshot = args.snapshot.resolve(strict=True) args.output.mkdir(parents=True, exist_ok=False) os.environ["HF_HUB_OFFLINE"] = "1" os.environ["TRANSFORMERS_OFFLINE"] = "1" logging.basicConfig(level=logging.INFO) # Check files and the installed patch before loading model weights. for command in ( [sys.executable, str(args.snapshot / "check_download.py"), str(args.snapshot), "--hash", "--runtime"], [sys.executable, str(args.snapshot / "compatibility/cache-tests/verify_server_patch.py")], ): subprocess.run(command, check=True) import mlx.core as mx from transformers import AutoTokenizer assert mx.metal.is_available(), "Metal is required; a CPU run cannot validate this crash" assert mx.default_device() != mx.cpu, "The validation must execute on the Metal GPU" assert importlib.metadata.version("mlx") == "0.32.2" assert importlib.metadata.version("mlx-lm") == "0.32.0" harness_path = args.snapshot / "compatibility/cache-tests/test_server_cache_http.py" spec = importlib.util.spec_from_file_location("cache_harness", harness_path) harness = importlib.util.module_from_spec(spec) spec.loader.exec_module(harness) tokenizer = AutoTokenizer.from_pretrained(args.snapshot, local_files_only=True) def prompt_length(messages): tokens = tokenizer.apply_chat_template( messages, tokenize=True, return_dict=False, add_generation_prompt=True, enable_thinking=False, ) assert isinstance(tokens, list) and tokens and isinstance(tokens[0], int) return len(tokens) def make_messages(count): return [ {"role": "system", "content": "Read the records. Reply briefly to the final question."}, {"role": "user", "content": ("Record: alpha beta gamma delta epsilon.\n" * count)}, {"role": "assistant", "content": "I have read the records."}, {"role": "user", "content": "Reply with the single word READY."}, ] lo, hi = 1, args.target_prompt_tokens while lo < hi: mid = (lo + hi) // 2 if prompt_length(make_messages(mid)) < args.target_prompt_tokens: lo = mid + 1 else: hi = mid messages = make_messages(lo) initial_tokens = prompt_length(messages) assert args.target_prompt_tokens <= initial_tokens < args.target_prompt_tokens + 128 print(f"Validated synthetic prompt length: {initial_tokens} tokens", flush=True) report = {"status": "RUNNING", "physical_memory_bytes": physical_bytes, "offline": True, "model_weights_changed": False, "memory_limit_overrides": False, "device": str(mx.default_device()), "device_info": mx.device_info(), "packages": {name: importlib.metadata.version(name) for name in ("mlx", "mlx-lm", "transformers", "huggingface_hub")}, "prompt_origin": "Synthetic repeated records; no third-party conversation", "initial_prompt_tokens": initial_tokens, "requests": []} output_file = args.output / "results.json" stopped = threading.Event() def save(): output_file.write_text(json.dumps(report, indent=2) + "\n") def monitor(generator): with (args.output / "memory.jsonl").open("w") as stream: while not stopped.is_set(): cache = generator.prompt_cache row = {"time": time.time(), "active_mlx_bytes": mx.get_active_memory(), "allocator_cache_bytes": mx.get_cache_memory(), "peak_mlx_bytes": mx.get_peak_memory(), "retained_bytes": cache.nbytes, "retained_limit": cache.max_bytes, "retained_sequences": len(cache), "worker_available": generator.generation_available(), "swap": subprocess.check_output(["sysctl", "-n", "vm.swapusage"], text=True).strip()} stream.write(json.dumps(row) + "\n") stream.flush() stopped.wait(10) save() try: with harness.running(args.snapshot, []) as (url, generator, cli): observer = threading.Thread(target=monitor, args=(generator,), daemon=True) observer.start() try: for turn in range(3): body = {"model": "default_model", "messages": messages, "max_tokens": 16, "temperature": 0, "chat_template_kwargs": {"enable_thinking": False}} req = urllib.request.Request(url + "/v1/chat/completions", data=json.dumps(body).encode(), headers={"Content-Type": "application/json"}) started = time.monotonic() with urllib.request.urlopen(req, timeout=3600) as response: result = json.load(response) assert response.status == 200 choice = result["choices"][0] content = choice["message"].get("content") assert isinstance(content, str) and content.strip(), result assert generator.generation_available() cache = generator.prompt_cache row = {"turn": turn + 1, "elapsed_seconds": time.monotonic() - started, "usage": result["usage"], "generation": content, "finish_reason": choice["finish_reason"], "retained_bytes": cache.nbytes, "retained_limit": cache.max_bytes, "retained_sequences": len(cache), "peak_mlx_bytes": mx.get_peak_memory()} report["requests"].append(row) save() print(json.dumps(row), flush=True) messages.extend([{"role": "assistant", "content": content}, {"role": "user", "content": "Reply with READY again."}]) hits = [row["usage"]["prompt_tokens_details"]["cached_tokens"] for row in report["requests"][1:]] report["warm_reuse_observed"] = any(hits) report["warm_reuse_fraction"] = [ row["usage"]["prompt_tokens_details"]["cached_tokens"] / row["usage"]["prompt_tokens"] for row in report["requests"][1:]] assert all(fraction >= 0.5 for fraction in report["warm_reuse_fraction"]), \ "Requests completed, but the cache did not reuse most of the long history" report["status"] = "PASS_FULL_MODEL_SYNTHETIC_LONG_CONTEXT_FOLLOWUPS" save() finally: stopped.set() observer.join(timeout=15) except BaseException as error: report["status"] = "FAIL" report["error"] = f"{type(error).__name__}: {error}" save() raise if __name__ == "__main__": main()