#!/usr/bin/env python3 # Copyright 2026 Modilify # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0 """Standalone text trial-inference CLI for native Modilify Mk1 MLX.""" from __future__ import annotations import argparse import sys from pathlib import Path import mlx.core as mx ROOT = Path(__file__).resolve().parent if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from modilify_mlx.generate import generate from modilify_mlx.modeling import load def _build_prompt_ids(model_path: Path, prompt: str, enable_thinking: bool): from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained(str(model_path), trust_remote_code=True) messages = [{"role": "user", "content": prompt}] token_ids = tokenizer.apply_chat_template( messages, tokenize=True, add_generation_prompt=True, enable_thinking=enable_thinking, ) if hasattr(token_ids, "input_ids"): token_ids = token_ids.input_ids elif isinstance(token_ids, dict): token_ids = token_ids["input_ids"] if hasattr(token_ids, "tolist"): token_ids = token_ids.tolist() if token_ids and isinstance(token_ids[0], (list, tuple)): token_ids = token_ids[0] return [int(token) for token in token_ids], tokenizer def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", type=Path, default=ROOT) parser.add_argument("--prompt", default="Explain why the sky is blue.") parser.add_argument("--max-new-tokens", type=int, default=128) parser.add_argument("--temperature", type=float, default=None) parser.add_argument("--enable-thinking", action="store_true") parser.add_argument("--seed", type=int, default=0) parser.add_argument( "--expert-bits", type=int, default=16, help="Quantize MoE experts to this bit width (16 keeps bf16).", ) parser.add_argument( "--profile", action="store_true", help="Print per-phase denoise timings.", ) args = parser.parse_args() print(f"[mk1] loading {args.model}", flush=True) model, config = load(args.model, expert_bits=args.expert_bits) if config.model_type != "modilify_mk1": raise SystemExit(f"Refusing model_type={config.model_type!r}") prompt_ids, tokenizer = _build_prompt_ids( args.model, args.prompt, args.enable_thinking ) print( f"[mk1] prompt_tokens={len(prompt_ids)} canvas={config.canvas_length} " f"temp={args.temperature if args.temperature is not None else config.denoise_temperature}", flush=True, ) profiler = None if args.profile: from modilify_mlx.profile import DenoiseProfiler profiler = DenoiseProfiler() output = generate( model, mx.array([prompt_ids], dtype=mx.int32), max_new_tokens=args.max_new_tokens, temperature=args.temperature, seed=args.seed, profiler=profiler, ) text = tokenizer.decode(output.generated_ids, skip_special_tokens=False) visible = tokenizer.decode(output.generated_ids, skip_special_tokens=True) print( f"[mk1] stop={output.stop_reason} denoise={output.denoise_steps} " f"committed={output.generated_length} avg_commit={output.average_commit_len:.2f} " f"tpf={output.tokens_per_forward:.2f} jumps={output.jump_count}", flush=True, ) print( f"[mk1] prefill={output.prefill_seconds:.3f}s " f"first_denoise={output.first_denoise_seconds:.3f}s " f"generate={output.generate_seconds:.3f}s " f"hd/s={output.heavy_denoise_per_second:.3f} " f"steady_hd/s={output.steady_heavy_denoise_per_second:.3f} " f"tok/s={output.tokens_per_second:.2f}", flush=True, ) if profiler is not None: print(profiler.summary(), flush=True) print("--- raw ---") print(text) print("--- visible ---") print(visible) if __name__ == "__main__": main()