Modilify-Mk1-MLX / generate_modilify.py
ydy9038074's picture
Publish Modilify Mk1 MLX runtime
a066584 verified
Raw
History Blame Contribute Delete
4.04 kB
#!/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()