File size: 4,041 Bytes
a066584
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
#!/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()