Text Generation
MLX
Safetensors
modilify_mk2
diffusion
mixture-of-experts
custom-code
modilify-mk2
conversational
Instructions to use modilify/Modilify-Mk2-preview-mlx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use modilify/Modilify-Mk2-preview-mlx with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("modilify/Modilify-Mk2-preview-mlx") prompt = "Write a story about Einstein" messages = [{"role": "user", "content": prompt}] prompt = tokenizer.apply_chat_template( messages, add_generation_prompt=True ) text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Pi
How to use modilify/Modilify-Mk2-preview-mlx with Pi:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure the model in Pi
# Install Pi: npm install -g @earendil-works/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "mlx-lm": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "modilify/Modilify-Mk2-preview-mlx" } ] } } }Run Pi
# Start Pi in your project directory: pi
- MLX LM
How to use modilify/Modilify-Mk2-preview-mlx with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Interactive chat REPL mlx_lm.chat --model "modilify/Modilify-Mk2-preview-mlx"
Run an OpenAI-compatible server
# Install MLX LM uv tool install mlx-lm # Start the server mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx" # Calling the OpenAI-compatible server with curl curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview-mlx", "messages": [ {"role": "user", "content": "Hello"} ] }' - Hermes Agent
How to use modilify/Modilify-Mk2-preview-mlx with Hermes Agent:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default modilify/Modilify-Mk2-preview-mlx
Run Hermes
hermes
- Atomic Chat
- OpenClaw
How to use modilify/Modilify-Mk2-preview-mlx with OpenClaw:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "modilify/Modilify-Mk2-preview-mlx" \ --custom-provider-id mlx-lm \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
Download inference.py from modilify/Modilify-Mk2-preview-mlx: direct link, hf CLI and curl.
- Browser
- Download file 56.5 kB
-
https://huggingface.co/modilify/Modilify-Mk2-preview-mlx/resolve/main/inference.py
- Command line
-
hf download hf://modilify/Modilify-Mk2-preview-mlx/inference.py
-
curl -L -o inference.py https://huggingface.co/modilify/Modilify-Mk2-preview-mlx/resolve/main/inference.py
56.5 kB
| """Native MLX schema25 inference with bounded KV reuse and continuous admission.""" | |
| from __future__ import annotations | |
| import argparse | |
| import copy | |
| import hashlib | |
| import json | |
| import math | |
| import queue | |
| import sys | |
| import threading | |
| import time | |
| from collections import OrderedDict, deque | |
| from dataclasses import dataclass, fields, replace | |
| from pathlib import Path | |
| from typing import Any, Iterator | |
| import mlx.core as mx | |
| from modilify_mk2.configuration_modilify_mk2 import DENOISE_TEMPERATURE | |
| from modilify_mk2.mlx_commit_policy import ( | |
| fused_commit_failure_rate, infer_commit_reason, select_commit_lengths, | |
| ) | |
| from modilify_mk2.mlx_model import MLXCanvasOutput, MLXModilifyMk2 | |
| from modilify_mk2.runtime import MLXRuntime, load_runtime | |
| from modilify_mk2.mlx_state import MLXLatentState, MLXRollingState | |
| from modilify_mk2.chat import apply_chat_template | |
| def parse_bool(value: str | bool) -> bool: | |
| """Parse explicit CLI booleans such as ``--think true``.""" | |
| if isinstance(value, bool): | |
| return value | |
| normalized = value.strip().lower() | |
| if normalized in {"1", "true", "yes", "on"}: | |
| return True | |
| if normalized in {"0", "false", "no", "off"}: | |
| return False | |
| raise argparse.ArgumentTypeError("expected true or false") | |
| REQUEST_FIELDS = frozenset( | |
| { | |
| "request_id", | |
| "prompt", | |
| "messages", | |
| "max_new_tokens", | |
| "max_denoising_steps", | |
| "seed", | |
| "think", | |
| } | |
| ) | |
| class ContinuousRequest: | |
| """One validated machine-mode request before chat-template tokenization.""" | |
| request_id: str | |
| messages: list[dict[str, Any]] | |
| max_new_tokens: int | |
| max_denoising_steps: int | None | |
| seed: int | |
| think: bool | |
| prompt: str | None = None | |
| def stable_request_seed(base_seed: int, request_id: str) -> int: | |
| """Derive a scheduler-independent non-negative seed from a stable request ID.""" | |
| payload = f"{base_seed}\0{request_id}".encode("utf-8") | |
| return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") & ((1 << 63) - 1) | |
| def _positive_int(record: dict[str, Any], name: str, default: int | None) -> int | None: | |
| value = record.get(name, default) | |
| if value is None: | |
| return default | |
| if isinstance(value, bool) or not isinstance(value, int) or value <= 0: | |
| raise ValueError(f"`{name}` must be a positive integer or null.") | |
| return value | |
| def parse_continuous_request( | |
| record: Any, | |
| *, | |
| default_max_new_tokens: int, | |
| default_max_denoising_steps: int | None, | |
| default_seed: int, | |
| default_think: bool, | |
| ) -> ContinuousRequest: | |
| """Validate one strict request object without silently inventing identity.""" | |
| if not isinstance(record, dict): | |
| raise ValueError("Each request line must be a JSON object.") | |
| unknown = sorted(set(record).difference(REQUEST_FIELDS)) | |
| if unknown: | |
| raise ValueError(f"Unknown request fields: {unknown}") | |
| request_id = record.get("request_id") | |
| if not isinstance(request_id, str) or not request_id.strip(): | |
| raise ValueError("`request_id` must be a non-empty string.") | |
| has_prompt = "prompt" in record | |
| has_messages = "messages" in record | |
| if has_prompt == has_messages: | |
| raise ValueError("Exactly one of `prompt` or `messages` is required.") | |
| prompt = None | |
| if has_prompt: | |
| prompt = record["prompt"] | |
| if not isinstance(prompt, str): | |
| raise ValueError("`prompt` must be a string.") | |
| messages = [{"role": "user", "content": prompt}] | |
| else: | |
| messages = record["messages"] | |
| if not isinstance(messages, list) or not messages: | |
| raise ValueError("`messages` must be a non-empty list.") | |
| if any( | |
| not isinstance(message, dict) | |
| or "role" not in message | |
| or "content" not in message | |
| for message in messages | |
| ): | |
| raise ValueError("Every message must be an object with `role` and `content`.") | |
| think = record.get("think", default_think) | |
| if not isinstance(think, bool): | |
| raise ValueError("`think` must be a boolean.") | |
| supplied_seed = record.get("seed") | |
| if supplied_seed is not None and ( | |
| isinstance(supplied_seed, bool) or not isinstance(supplied_seed, int) | |
| ): | |
| raise ValueError("`seed` must be an integer or null.") | |
| seed = ( | |
| stable_request_seed(default_seed, request_id) | |
| if supplied_seed is None | |
| else supplied_seed | |
| ) | |
| return ContinuousRequest( | |
| request_id=request_id, | |
| prompt=prompt, | |
| messages=messages, | |
| max_new_tokens=int( | |
| _positive_int(record, "max_new_tokens", default_max_new_tokens) | |
| ), | |
| max_denoising_steps=_positive_int( | |
| record, "max_denoising_steps", default_max_denoising_steps | |
| ), | |
| seed=seed, | |
| think=think, | |
| ) | |
| def _extract_input_ids(encoded: Any) -> list[int]: | |
| value = encoded.get("input_ids") if isinstance(encoded, dict) else encoded | |
| if hasattr(encoded, "input_ids"): | |
| value = encoded.input_ids | |
| if isinstance(value, tuple): | |
| value = list(value) | |
| if isinstance(value, list) and len(value) == 1 and isinstance(value[0], list): | |
| value = value[0] | |
| if not isinstance(value, list) or any( | |
| isinstance(token, bool) or not isinstance(token, int) for token in value | |
| ): | |
| raise ValueError("Chat template did not return one integer token sequence.") | |
| if not value: | |
| raise ValueError("Chat template returned an empty prompt.") | |
| return [int(token) for token in value] | |
| def _seed(value: int, salt: int) -> mx.array: | |
| return mx.random.key((int(value) + int(salt)) % (2**32)) | |
| def _logical(value: mx.array, head: mx.array) -> mx.array: | |
| canvas = value.shape[1] | |
| index = (head[:, None] + mx.arange(canvas)[None, :]) % canvas | |
| return mx.take_along_axis(value, index, axis=1) | |
| def _physical(value: mx.array, head: mx.array) -> mx.array: | |
| canvas = value.shape[1] | |
| index = (mx.arange(canvas)[None, :] - head[:, None]) % canvas | |
| return mx.take_along_axis(value, index, axis=1) | |
| def _concat_rows(values: list[Any]) -> Any: | |
| first = values[0] | |
| if isinstance(first, mx.array): | |
| return mx.concatenate(values, axis=0) | |
| return type(first)(**{ | |
| item.name: _concat_rows([getattr(value, item.name) for value in values]) | |
| for item in fields(first) | |
| }) | |
| def _slice_row(value: Any, row: int) -> Any: | |
| if isinstance(value, mx.array): | |
| return value[row:row + 1] | |
| return type(value)(**{ | |
| item.name: _slice_row(getattr(value, item.name), row) | |
| for item in fields(value) | |
| }) | |
| def _arrays(value: Any) -> list[mx.array]: | |
| if isinstance(value, mx.array): | |
| return [value] | |
| return [array for item in fields(value) | |
| for array in _arrays(getattr(value, item.name))] | |
| def _clone_cache(cache: list[Any], *, compact: bool = False) -> list[Any]: | |
| copied = [] | |
| arrays = [] | |
| for layer in cache: | |
| clone = copy.copy(layer) | |
| source = layer.state if compact else (layer.keys, layer.values) | |
| if source[0] is not None: | |
| clone.keys = mx.array(source[0]) | |
| clone.values = mx.array(source[1]) | |
| arrays.extend((clone.keys, clone.values)) | |
| copied.append(clone) | |
| if arrays: | |
| mx.eval(*arrays) | |
| return copied | |
| class _PrefixEntry: | |
| cache: list[Any] | |
| bytes: int | |
| class PrefixKVCache: | |
| """LRU cache of immutable, block-boundary prompt KV snapshots.""" | |
| def __init__(self, limit_bytes: int): | |
| self.limit_bytes = limit_bytes | |
| self.entries: OrderedDict[tuple[int, ...], _PrefixEntry] = OrderedDict() | |
| self.bytes = 0 | |
| self.hits = 0 | |
| self.reused_tokens = 0 | |
| def longest(self, ids: tuple[int, ...]) -> tuple[int, list[Any] | None]: | |
| lengths = sorted( | |
| {len(key) for key in self.entries if len(key) <= len(ids)}, | |
| reverse=True, | |
| ) | |
| for length in lengths: | |
| key = ids[:length] | |
| entry = self.entries.get(key) | |
| if entry is not None: | |
| self.entries.move_to_end(key) | |
| self.hits += 1 | |
| self.reused_tokens += length | |
| return length, _clone_cache(entry.cache) | |
| return 0, None | |
| def put(self, ids: tuple[int, ...], cache: list[Any]) -> None: | |
| if self.limit_bytes <= 0 or ids in self.entries: | |
| return | |
| size = sum(int(value.nbytes) for layer in cache | |
| for value in (layer.state if layer.keys is not None else ()) | |
| if value is not None) | |
| if size > self.limit_bytes: | |
| return | |
| snapshot = _clone_cache(cache, compact=True) | |
| while self.entries and self.bytes + size > self.limit_bytes: | |
| _, victim = self.entries.popitem(last=False) | |
| self.bytes -= victim.bytes | |
| self.entries[ids] = _PrefixEntry(snapshot, size) | |
| self.bytes += size | |
| class InferenceRow: | |
| request: ContinuousRequest | |
| prompt_ids: tuple[int, ...] | |
| cache: list[Any] | |
| rolling: MLXRollingState | |
| generated: list[int] | |
| repetition_seen: set[int] | |
| created: float | |
| admitted: float | |
| denoise_steps: int = 0 | |
| jumps: int = 0 | |
| shifts: int = 0 | |
| first_token_at: float | None = None | |
| stop_reason: str | None = None | |
| first_denoise_at: float | None = None | |
| last_denoise_at: float | None = None | |
| prefill_seconds: float = 0.0 | |
| class PrefillRow: | |
| request: ContinuousRequest | |
| prompt_ids: tuple[int, ...] | |
| cache: list[Any] | |
| offset: int | |
| reused: int | |
| created: float | |
| seconds: float = 0.0 | |
| class _DenoiseWork: | |
| rows: list[InferenceRow] | |
| rolling: MLXRollingState | |
| output: Any | |
| proposal: mx.array | |
| remaining: mx.array | |
| physical_positions: mx.array | |
| next_latent: MLXLatentState | |
| policy: Any | |
| def _forward_independent_rows(model: MLXModilifyMk2, | |
| rows: list[InferenceRow]) -> list[MLXCanvasOutput]: | |
| """Interleave decoder layers while preserving independent singleton rows.""" | |
| decoder = model.model.decoder | |
| latent = model.latent_deliberation | |
| working_bus, persistent_bus = latent.working_memory_bus, latent.persistent_memory_bus | |
| contexts = [] | |
| for row in rows: | |
| rolling = row.rolling | |
| canvas, state, head = rolling.canvas, rolling.latent, rolling.head | |
| batch, length = canvas.shape | |
| tokens = decoder.embed_tokens(canvas) * decoder.embed_scale | |
| working, next_state = latent(token_embeddings=tokens, | |
| confidence=state.confidence, | |
| entropy=state.entropy, state=state, | |
| canvas_head=head) | |
| physical = (head[:, None] + mx.arange(length)[None, :]) % length | |
| gather = mx.broadcast_to(physical[:, :, None], tokens.shape) | |
| logical_tokens = mx.take_along_axis(tokens, mx.stop_gradient(gather), axis=1) | |
| logical_working = mx.take_along_axis(working, mx.stop_gradient(gather), axis=1) | |
| hidden = model._merge_context(logical_tokens, logical_working) | |
| prefix_length = int(row.cache[0].offset) | |
| full_mask = mx.ones((batch, prefix_length + length), mx.bool_) | |
| masks = decoder._make_decoder_masks(hidden, row.cache, full_mask) | |
| seen = latent.logical_seen(state, head) | |
| contexts.append({"hidden": hidden, "working": working, "state": next_state, | |
| "tokens": tokens, "head": head, "masks": masks, | |
| "offset": prefix_length, | |
| "working_kv": working_bus.prepare_kv((logical_working, seen)), | |
| "persistent_kv": persistent_bus.prepare_kv(state.memory_slots, seen)}) | |
| reader = 0 | |
| for index, layer in enumerate(decoder.layers): | |
| for row, context in zip(rows, contexts, strict=True): | |
| hidden = layer(context["hidden"], context["masks"][layer.layer_type], | |
| row.cache[index], decoder=True, offset=context["offset"]) | |
| if layer.layer_type == "full_attention": | |
| if context["working_kv"] is not None and reader < working_bus.num_readers: | |
| hidden = working_bus.read(hidden, reader, *context["working_kv"]) | |
| if context["persistent_kv"] is not None and reader < persistent_bus.num_readers: | |
| hidden = persistent_bus.read(hidden, reader, *context["persistent_kv"]) | |
| context["hidden"] = hidden | |
| # Explicit layer boundaries keep the execution order interleaved and | |
| # bound intermediate activations to the configured pipeline depth. | |
| mx.async_eval(*(context["hidden"] for context in contexts)) | |
| if layer.layer_type == "full_attention": | |
| reader += 1 | |
| outputs = [] | |
| for context in contexts: | |
| hidden = decoder.norm(context["hidden"]) | |
| length = hidden.shape[1] | |
| inverse = (mx.arange(length)[None, :] - context["head"][:, None]) % length | |
| heavy = mx.take_along_axis(hidden, mx.stop_gradient(mx.broadcast_to( | |
| inverse[:, :, None], hidden.shape)), axis=1) | |
| outputs.append(MLXCanvasOutput(heavy, context["working"], | |
| context["state"], context["tokens"])) | |
| return outputs | |
| def _noise(seed: int, step: int, canvas: int, vocab: int) -> mx.array: | |
| return mx.random.randint(0, vocab, (1, canvas), | |
| key=_seed(seed, 0x51A7 + step * 1000003)) | |
| def _empty_rolling(config: Any, seed: int, max_new_tokens: int, | |
| pad_token_id: int) -> MLXRollingState: | |
| canvas = int(config.canvas_length) | |
| vocab = int(config.text_config.vocab_size) | |
| latent = MLXLatentState.empty(1, canvas, int(config.latent_memory_slots), | |
| int(config.latent_dim), enable_gdn2=True) | |
| latent = replace(latent, entropy=mx.full((1, canvas), math.log(vocab), mx.float32)) | |
| initial = mx.where(mx.arange(canvas)[None, :] < max_new_tokens, | |
| _noise(seed, 0, canvas, vocab), pad_token_id) | |
| return MLXRollingState(initial, latent, | |
| mx.zeros((1,), mx.int32)) | |
| def _inference_statistics(hidden: mx.array, weight: mx.array, | |
| rows: list[InferenceRow], *, softcap: float, | |
| chunk_size: int, penalty: float, | |
| excluded: set[int], | |
| top_k: int | None = 40, | |
| min_p: float | None = 0.05) -> tuple[mx.array, ...]: | |
| """Exact chunked Top-K and Min-P Gumbel sampling without a retained full-vocabulary matrix.""" | |
| batch, canvas, dim = hidden.shape | |
| flat = hidden.reshape(batch * canvas, dim) | |
| neg_inf = mx.full((batch * canvas,), -mx.inf, mx.float32) | |
| log_z = neg_inf | |
| best_gumbel = neg_inf | |
| chosen_score = mx.zeros_like(log_z) | |
| chosen = mx.zeros((batch * canvas,), mx.int32) | |
| greedy_score = neg_inf | |
| greedy = mx.zeros_like(chosen) | |
| moment_max = neg_inf | |
| moment_sum = mx.zeros_like(log_z) | |
| moment_weighted = mx.zeros_like(log_z) | |
| repetition_mask = None | |
| if penalty != 1.0: | |
| masks = [] | |
| for row in rows: | |
| mask = mx.zeros((weight.shape[0],), mx.bool_) | |
| eligible = sorted(row.repetition_seen - excluded) | |
| if eligible: | |
| mask[mx.array(eligible, mx.int32)] = True | |
| masks.append(mask) | |
| repetition_mask = mx.stack(masks, axis=0) | |
| use_constrained = top_k is not None and top_k > 0 | |
| chunk_cand_scores = [] | |
| chunk_cand_tokens = [] | |
| for chunk_number, start in enumerate(range(0, weight.shape[0], chunk_size)): | |
| stop = min(start + chunk_size, weight.shape[0]) | |
| score = mx.tanh((flat @ weight[start:stop].T).astype(mx.float32) / softcap) * softcap | |
| if penalty != 1.0: | |
| assert repetition_mask is not None | |
| row_scores = score.reshape(batch, canvas, stop - start) | |
| seen = repetition_mask[:, None, start:stop] | |
| changed = mx.where(row_scores < 0, row_scores * penalty, row_scores / penalty) | |
| score = mx.where(seen, changed, row_scores).reshape(batch * canvas, stop - start) | |
| score = score / DENOISE_TEMPERATURE | |
| log_z = mx.logaddexp(log_z, mx.logsumexp(score, axis=-1)) | |
| local_max = mx.max(score, axis=-1) | |
| better_greedy = local_max > greedy_score | |
| greedy = mx.where(better_greedy, mx.argmax(score, axis=-1).astype(mx.int32) + start, greedy) | |
| greedy_score = mx.maximum(greedy_score, local_max) | |
| shifted = mx.exp(score - local_max[:, None]) | |
| next_max = mx.maximum(moment_max, local_max) | |
| old_scale = mx.exp(moment_max - next_max) | |
| new_scale = mx.exp(local_max - next_max) | |
| moment_sum = moment_sum * old_scale + mx.sum(shifted, axis=-1) * new_scale | |
| moment_weighted = moment_weighted * old_scale + mx.sum(shifted * score, axis=-1) * new_scale | |
| moment_max = next_max | |
| if use_constrained: | |
| k = min(top_k, stop - start) | |
| chunk_top_idx = mx.stop_gradient(mx.argpartition(-score, kth=k - 1, axis=-1)[:, :k]).astype(mx.int32) | |
| chunk_top_score = mx.take_along_axis(score, chunk_top_idx, axis=-1) | |
| chunk_cand_scores.append(chunk_top_score) | |
| chunk_cand_tokens.append(chunk_top_idx + start) | |
| else: | |
| uniform = mx.concatenate([ | |
| mx.random.uniform(shape=(canvas, stop - start), | |
| key=_seed(row.request.seed, | |
| 0xC09A + row.denoise_steps * 1000003 + chunk_number)) | |
| for row in rows | |
| ], axis=0) | |
| uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7) | |
| gumbel = score - mx.log(-mx.log(uniform)) | |
| local_index = mx.argmax(gumbel, axis=-1) | |
| local_best = mx.max(gumbel, axis=-1) | |
| better = local_best > best_gumbel | |
| chosen_score = mx.where(better, mx.take_along_axis(score, local_index[:, None], axis=-1)[:, 0], chosen_score) | |
| chosen = mx.where(better, local_index + start, chosen) | |
| best_gumbel = mx.maximum(best_gumbel, local_best) | |
| # Materialize each chunk without a CPU/GPU barrier. Include candidates | |
| # so their argpartition graph does not retain every vocabulary slab. | |
| live = [log_z, greedy, greedy_score, moment_max, moment_sum, moment_weighted] | |
| if use_constrained: | |
| live.extend((chunk_cand_scores[-1], chunk_cand_tokens[-1])) | |
| else: | |
| live.extend((best_gumbel, chosen_score, chosen)) | |
| mx.async_eval(*live) | |
| if use_constrained: | |
| all_cand_scores = mx.concatenate(chunk_cand_scores, axis=-1) | |
| all_cand_tokens = mx.concatenate(chunk_cand_tokens, axis=-1) | |
| global_k = min(top_k, all_cand_scores.shape[-1]) | |
| global_idx = mx.stop_gradient(mx.argpartition(-all_cand_scores, kth=global_k - 1, axis=-1)[:, :global_k]).astype(mx.int32) | |
| cand_scores = mx.take_along_axis(all_cand_scores, global_idx, axis=-1) | |
| cand_tokens = mx.take_along_axis(all_cand_tokens, global_idx, axis=-1) | |
| if min_p is not None and min_p > 0.0: | |
| thresh = greedy_score + math.log(float(min_p)) | |
| valid_cand = cand_scores >= thresh[:, None] | |
| eligible_scores = mx.where(valid_cand, cand_scores, -mx.inf) | |
| else: | |
| eligible_scores = cand_scores | |
| uniform = mx.concatenate([ | |
| mx.random.uniform(shape=(canvas, global_k), | |
| key=_seed(row.request.seed, | |
| 0xC09A + row.denoise_steps * 1000003)) | |
| for row in rows | |
| ], axis=0) | |
| uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7) | |
| gumbel = eligible_scores - mx.log(-mx.log(uniform)) | |
| chosen_local = mx.stop_gradient(mx.argmax(gumbel, axis=-1)) | |
| chosen_score = mx.take_along_axis(cand_scores, chosen_local[:, None], axis=-1)[:, 0] | |
| chosen = mx.take_along_axis(cand_tokens, chosen_local[:, None], axis=-1)[:, 0] | |
| confidence = mx.clip(mx.exp(chosen_score - log_z), 0.0, 1.0) | |
| greedy_confidence = mx.clip(mx.exp(greedy_score - log_z), 0.0, 1.0) | |
| entropy = log_z - moment_weighted / mx.maximum(moment_sum, 1.17549435e-38) | |
| shape = (batch, canvas) | |
| return (chosen.reshape(shape), confidence.reshape(shape), entropy.reshape(shape), | |
| greedy.reshape(shape), greedy_confidence.reshape(shape)) | |
| class MLXContinuousEngine: | |
| def __init__(self, runtime: MLXRuntime, *, prefix_mib: int, | |
| prefill_chunk: int, vocab_chunk: int, max_batch_rows: int, | |
| max_batch_tokens: int, emit: Any, pipeline_depth: int = 2, | |
| token_events: bool = True): | |
| self.runtime = runtime | |
| if min(prefill_chunk, vocab_chunk, max_batch_rows, pipeline_depth) <= 0: | |
| raise ValueError("Chunk sizes, batch rows and pipeline depth must be positive.") | |
| if max_batch_tokens < int(runtime.config.canvas_length): | |
| raise ValueError("--max-batch-tokens must fit at least one canvas.") | |
| self.prefix = PrefixKVCache(prefix_mib * 1024 * 1024) | |
| self.prefill_chunk = prefill_chunk | |
| self.vocab_chunk = vocab_chunk | |
| self.max_batch_rows = max_batch_rows | |
| self.max_batch_tokens = max_batch_tokens | |
| self.emit = emit | |
| self.token_events = token_events | |
| self.pipeline_depth = pipeline_depth | |
| self.forward_count = 0 | |
| self.active_row_steps = 0 | |
| self.denoise_seconds = 0.0 | |
| self.max_inflight = 0 | |
| self.peak_active_requests = 0 | |
| generation = runtime.generation | |
| pad_token_id = generation.pad_token_id | |
| if pad_token_id is None: | |
| pad_token_id = getattr(runtime.config, "pad_token_id", None) | |
| if isinstance(pad_token_id, (list, tuple)): | |
| pad_token_id = pad_token_id[0] | |
| self.pad_token_id = int(0 if pad_token_id is None else pad_token_id) | |
| configured_eos = generation.eos_token_id or runtime.config.eos_token_id | |
| if isinstance(configured_eos, int): | |
| configured_eos = [configured_eos] | |
| self.turn_end = int(runtime.config.turn_end_token_id if | |
| generation.turn_end_token_id is None else | |
| generation.turn_end_token_id) | |
| self.stops = tuple(dict.fromkeys((self.turn_end, *(int(x) for x in configured_eos or ())))) | |
| token_values = [generation.pad_token_id, generation.bos_token_id, | |
| generation.eos_token_id, generation.turn_end_token_id, | |
| getattr(runtime.config, "image_token_id", None), | |
| generation.repetition_penalty_exclude_token_ids] | |
| self.excluded = set() | |
| for value in token_values: | |
| if isinstance(value, int): | |
| self.excluded.add(int(value)) | |
| elif isinstance(value, (tuple, list, set)): | |
| self.excluded.update(int(x) for x in value if x is not None) | |
| def begin_prefill(self, request: ContinuousRequest, created: float) -> PrefillRow: | |
| encoded = apply_chat_template(self.runtime.tokenizer, request.messages, | |
| think=request.think) | |
| prompt = tuple(_extract_input_ids(encoded)) | |
| if len(prompt) + request.max_new_tokens > int(self.runtime.config.text_config.max_position_embeddings): | |
| raise ValueError("Prompt plus response exceeds the model position limit.") | |
| prefix_length, cache = self.prefix.longest(prompt) | |
| if cache is None: | |
| cache = self.runtime.model.model.encoder.make_cache() | |
| return PrefillRow(request, prompt, cache, prefix_length, prefix_length, created) | |
| def prefill_step(self, work: PrefillRow) -> InferenceRow | None: | |
| """Execute at most one prefill chunk before yielding to active rows.""" | |
| started = time.perf_counter() | |
| if work.offset < len(work.prompt_ids): | |
| end = min(work.offset + self.prefill_chunk, len(work.prompt_ids)) | |
| block = mx.array(work.prompt_ids[work.offset:end], mx.int32)[None, :] | |
| _, work.cache = self.runtime.model.model.encoder(block, cache=work.cache) | |
| mx.eval(*(value for layer in work.cache for value in (layer.keys, layer.values) | |
| if isinstance(value, mx.array))) | |
| work.offset = end | |
| self.prefix.put(work.prompt_ids[:end], work.cache) | |
| if work.offset < len(work.prompt_ids): | |
| work.seconds += time.perf_counter() - started | |
| return None | |
| request, prompt = work.request, work.prompt_ids | |
| cache = _clone_cache(work.cache, compact=True) | |
| work.seconds += time.perf_counter() - started | |
| seen = {token for token in prompt if token not in self.excluded} | |
| row = InferenceRow(request, prompt, cache, | |
| _empty_rolling(self.runtime.config, request.seed, | |
| request.max_new_tokens, self.pad_token_id), [], | |
| seen, work.created, time.perf_counter(), | |
| prefill_seconds=work.seconds) | |
| self.emit({"event": "request_started", "request_id": request.request_id, | |
| "prompt_tokens": len(prompt), "prefix_cache_tokens": work.reused, | |
| "prefill_seconds": row.prefill_seconds}) | |
| return row | |
| def admit(self, request: ContinuousRequest, created: float) -> InferenceRow: | |
| work = self.begin_prefill(request, created) | |
| while True: | |
| row = self.prefill_step(work) | |
| if row is not None: | |
| return row | |
| def _append_encoder(self, rows: list[InferenceRow], blocks: list[list[int]]) -> None: | |
| # Keep the encoder's GEMM and expert routing shapes independent of the | |
| # cohort too. Only the scheduler and GPU submission are concurrent. | |
| for row, tokens in zip(rows, blocks, strict=True): | |
| if tokens: | |
| ids = mx.array(tokens, mx.int32)[None, :] | |
| _, row.cache = self.runtime.model.model.encoder(ids, cache=row.cache) | |
| mx.async_eval(*(value for layer in row.cache | |
| for value in (layer.keys, layer.values) | |
| if isinstance(value, mx.array))) | |
| def step(self, rows: list[InferenceRow]) -> list[InferenceRow]: | |
| if not rows: | |
| return [] | |
| if len(rows) > min(self.max_batch_rows, self.max_batch_tokens // | |
| int(self.runtime.config.canvas_length)): | |
| raise ValueError("Denoise cohort exceeds the configured row/token budget.") | |
| if len({id(row) for row in rows}) != len(rows): | |
| raise ValueError("A request may appear only once in a denoise cohort.") | |
| started = time.perf_counter() | |
| finished = [] | |
| for start in range(0, len(rows), self.pipeline_depth): | |
| cohort = rows[start:start + self.pipeline_depth] | |
| outputs = _forward_independent_rows(self.runtime.model, cohort) | |
| pending = [self._prepare_step([row], output=output) | |
| for row, output in zip(cohort, outputs, strict=True)] | |
| self.max_inflight = max(self.max_inflight, len(pending)) | |
| for work in pending: | |
| finished.extend(self._finish_step(work)) | |
| self.denoise_seconds += time.perf_counter() - started | |
| return finished | |
| def _prepare_step(self, rows: list[InferenceRow], *, | |
| output: MLXCanvasOutput) -> _DenoiseWork: | |
| model, config, generation = (self.runtime.model, self.runtime.config, | |
| self.runtime.generation) | |
| rolling = _concat_rows([row.rolling for row in rows]) | |
| batch, canvas = rolling.canvas.shape | |
| stats = _inference_statistics( | |
| output.heavy_hidden, model.model.decoder.embed_tokens.weight, | |
| rows, softcap=float(config.text_config.final_logit_softcapping), | |
| chunk_size=self.vocab_chunk, penalty=float(generation.repetition_penalty), | |
| excluded=self.excluded, | |
| top_k=getattr(config, "commit_top_k", 40), | |
| min_p=getattr(config, "commit_min_p", 0.05), | |
| ) | |
| proposal, confidence, entropy, greedy, greedy_confidence = stats | |
| confidence = confidence.astype(mx.float32) | |
| entropy = entropy.astype(mx.float32) | |
| changed = (proposal != rolling.canvas).astype(mx.float32) | |
| remaining = mx.array([row.request.max_new_tokens - len(row.generated) | |
| for row in rows], mx.int32) | |
| physical_positions = (mx.arange(canvas)[None, :] - rolling.head[:, None]) % canvas | |
| valid = physical_positions < remaining[:, None] | |
| next_latent = replace( | |
| output.next_latent_state, confidence=confidence, entropy=entropy, | |
| age=rolling.latent.age + 1, token_changed=changed, | |
| confidence_delta=confidence - rolling.latent.confidence, | |
| entropy_delta=entropy - rolling.latent.entropy, | |
| ) | |
| next_latent = _concat_rows([ | |
| model.latent_deliberation.observe_state( | |
| _slice_row(next_latent, index), output.heavy_hidden[index:index + 1], | |
| output.working_state[index:index + 1], valid[index:index + 1], | |
| rolling.head[index:index + 1]) for index in range(batch) | |
| ]) | |
| policy = select_commit_lengths( | |
| _logical(proposal, rolling.head), | |
| _logical(fused_commit_failure_rate( | |
| confidence, entropy, entropy_weight=config.commit_entropy_weight, | |
| confidence_power=config.commit_confidence_power, | |
| top_k=getattr(config, "commit_top_k", None), | |
| min_p=getattr(config, "commit_min_p", None), | |
| target_confidence=getattr(config, "commit_target_confidence", None), | |
| failure_budget=float(config.commit_failure_budget)), rolling.head), | |
| _logical(fused_commit_failure_rate( | |
| rolling.latent.confidence, rolling.latent.entropy, | |
| entropy_weight=config.commit_entropy_weight, | |
| confidence_power=config.commit_confidence_power, | |
| top_k=getattr(config, "commit_top_k", None), | |
| min_p=getattr(config, "commit_min_p", None), | |
| target_confidence=getattr(config, "commit_target_confidence", None), | |
| failure_budget=float(config.commit_failure_budget)), rolling.head), | |
| _logical(greedy, rolling.head), | |
| _logical(fused_commit_failure_rate( | |
| greedy_confidence, entropy, entropy_weight=config.commit_entropy_weight, | |
| confidence_power=config.commit_confidence_power, | |
| top_k=getattr(config, "commit_top_k", None), | |
| min_p=getattr(config, "commit_min_p", None), | |
| target_confidence=getattr(config, "commit_target_confidence", None), | |
| failure_budget=float(config.commit_failure_budget)), rolling.head), | |
| ponder_steps=rolling.latent.ponder_steps, | |
| stagnation_steps=rolling.latent.stagnation_steps, | |
| active_rows=mx.ones((batch,), mx.bool_), remaining_lengths=remaining, | |
| failure_budget=float(config.commit_failure_budget), | |
| stop_token_id=self.stops, | |
| stagnation_threshold=int(generation.jump_on_no_progress_after), | |
| min_progress=float(generation.min_trajectory_progress), | |
| max_ponder_steps=int(generation.max_ponder_steps), | |
| valid_mask=mx.arange(canvas)[None, :] < remaining[:, None], | |
| ) | |
| mx.async_eval(policy.commit_lengths, policy.commit_token_ids, | |
| policy.jump_rows, policy.ponder_steps, policy.stagnation_steps, | |
| *_arrays(next_latent), output.heavy_hidden, output.working_state) | |
| self.forward_count += 1 | |
| return _DenoiseWork(rows, rolling, output, proposal, remaining, | |
| physical_positions, next_latent, policy) | |
| def _finish_step(self, work: _DenoiseWork) -> list[InferenceRow]: | |
| rows, rolling, output = work.rows, work.rolling, work.output | |
| proposal, remaining = work.proposal, work.remaining | |
| physical_positions, next_latent, policy = ( | |
| work.physical_positions, work.next_latent, work.policy) | |
| model, config, generation = (self.runtime.model, self.runtime.config, | |
| self.runtime.generation) | |
| batch, canvas = rolling.canvas.shape | |
| lengths = policy.commit_lengths.astype(mx.int32) | |
| selected = policy.commit_token_ids | |
| mx.eval(lengths, selected, policy.jump_rows) | |
| completed = time.perf_counter() | |
| for row in rows: | |
| row.last_denoise_at = completed | |
| if row.first_denoise_at is None: | |
| row.first_denoise_at = completed | |
| self.emit({"event": "first_denoise", "request_id": row.request.request_id, | |
| "ttfd_seconds": completed - row.created}) | |
| lengths_host = [int(x) for x in lengths.tolist()] | |
| maximum = max(lengths_host) | |
| selected_host = selected.tolist() | |
| blocks = [[int(x) for x in selected_host[index][:lengths_host[index]]] | |
| for index in range(batch)] | |
| # The commit policy has fixed these tokens. Stream them before the | |
| # writer and encoder append, which are only needed by the next step. | |
| for index, row in enumerate(rows): | |
| block = blocks[index] | |
| if not block: | |
| continue | |
| if row.first_token_at is None: | |
| row.first_token_at = time.perf_counter() | |
| row.generated.extend(block) | |
| row.repetition_seen.update(token for token in block | |
| if token not in self.excluded) | |
| if self.token_events: | |
| self.emit({"event": "token", "request_id": row.request.request_id, | |
| "token_ids": block, "generated_tokens": len(row.generated), | |
| "text": self.runtime.tokenizer.decode( | |
| row.generated, skip_special_tokens=False), | |
| "denoise_steps": row.denoise_steps + 1}) | |
| next_canvas = mx.where( | |
| (physical_positions < lengths[:, None]) & policy.jump_rows[:, None], | |
| _physical(selected, rolling.head), proposal, | |
| ) | |
| next_latent = replace(next_latent, ponder_steps=policy.ponder_steps, | |
| stagnation_steps=policy.stagnation_steps) | |
| if maximum: | |
| embeddings = model.model.decoder.embed_tokens(selected[:, :maximum]) | |
| embeddings = embeddings * model.model.decoder.embed_scale | |
| memory, _ = model.latent_deliberation.commit_write( | |
| memory=next_latent.memory_slots, | |
| working_state=output.working_state, | |
| heavy_hidden=output.heavy_hidden, | |
| committed_token_embeddings=embeddings, | |
| commit_lengths=lengths, | |
| prefix_lengths=mx.array([int(row.cache[0].offset) for row in rows], mx.int32), | |
| commit_reason=infer_commit_reason( | |
| lengths, jump_rows=policy.jump_rows, | |
| commit_token_ids=selected, terminal_token_ids=self.stops, | |
| ), | |
| canvas_head=rolling.head, max_commit=maximum, | |
| ) | |
| next_latent = replace( | |
| next_latent, memory_slots=memory, | |
| gdn2=replace(next_latent.gdn2, persistent=memory), | |
| ) | |
| self._append_encoder(rows, blocks) | |
| noise = mx.concatenate([ | |
| _noise(row.request.seed, row.denoise_steps + 1, canvas, | |
| int(config.text_config.vocab_size)) for row in rows | |
| ], axis=0) | |
| next_rolling = MLXRollingState(next_canvas, next_latent, | |
| rolling.head).advance_ring( | |
| lengths, noise, entropy_fill_value=math.log(config.text_config.vocab_size), | |
| ) | |
| logical_positions = (mx.arange(canvas)[None, :] - next_rolling.head[:, None]) % canvas | |
| newly_exposed = logical_positions >= (canvas - lengths)[:, None] | |
| next_rolling = replace( | |
| next_rolling, | |
| canvas=mx.where(newly_exposed & | |
| (logical_positions >= (remaining - lengths)[:, None]), | |
| self.pad_token_id, next_rolling.canvas), | |
| ) | |
| # The following step depends on these arrays on the same MLX stream. | |
| # Submit now and overlap the writer/refill with host stop/output work. | |
| mx.async_eval(*_arrays(next_rolling)) | |
| finished = [] | |
| jumped = policy.jump_rows.tolist() | |
| for index, row in enumerate(rows): | |
| row.rolling = _slice_row(next_rolling, index) | |
| row.denoise_steps += 1 | |
| self.active_row_steps += 1 | |
| row.jumps += int(jumped[index]) | |
| row.shifts += int(bool(blocks[index])) | |
| if self.turn_end in blocks[index]: | |
| row.stop_reason = "turn_end" | |
| elif any(token in self.stops for token in blocks[index]): | |
| row.stop_reason = "eos" | |
| elif len(row.generated) >= row.request.max_new_tokens: | |
| row.stop_reason = "max_new_tokens" | |
| elif (row.request.max_denoising_steps is not None and | |
| row.denoise_steps >= row.request.max_denoising_steps): | |
| row.stop_reason = "max_denoising_steps" | |
| else: | |
| bound = row.request.max_new_tokens * int(generation.max_ponder_steps) | |
| if row.denoise_steps >= bound: | |
| row.stop_reason = "episode_watchdog" | |
| if row.stop_reason is not None: | |
| finished.append(row) | |
| return finished | |
| def result(self, row: InferenceRow) -> dict[str, Any]: | |
| now = time.perf_counter() | |
| return {"event": "generation_result", "request_id": row.request.request_id, | |
| "status": "finished", "checkpoint": str(self.runtime.checkpoint), | |
| "step": self.runtime.step, "thinking": "on" if row.request.think else "off", | |
| "stop_reason": row.stop_reason, "prompt_tokens": len(row.prompt_ids), | |
| "generated_tokens": len(row.generated), "token_ids": row.generated, | |
| "text": self.runtime.tokenizer.decode(row.generated, skip_special_tokens=True), | |
| "denoise_steps": row.denoise_steps, "jump_count": row.jumps, | |
| "state_shift_count": row.shifts, | |
| "tokens_per_forward": len(row.generated) / max(row.denoise_steps, 1), | |
| "queue_seconds": row.admitted - row.created, | |
| "prefill_seconds": row.prefill_seconds, | |
| "ttfd_seconds": None if row.first_denoise_at is None else row.first_denoise_at - row.created, | |
| "denoise_steps_per_second": row.denoise_steps / max( | |
| (row.last_denoise_at or now) - row.admitted, 1e-9), | |
| "ttft_seconds": None if row.first_token_at is None else row.first_token_at - row.created, | |
| "elapsed_seconds": now - row.created, | |
| "prefix_cache_hits": self.prefix.hits, | |
| "prefix_cache_reused_tokens": self.prefix.reused_tokens, | |
| "mlx_active_memory_bytes": int(mx.get_active_memory()), | |
| "mlx_peak_memory_bytes": int(mx.get_peak_memory())} | |
| def iter_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any], | |
| max_queue_size: int, *, | |
| continue_on_error: bool = False) -> Iterator[dict[str, Any]]: | |
| """Bounded continuous admission with a prefill token budget per decode turn.""" | |
| active: deque[InferenceRow] = deque() | |
| prefilling: deque[PrefillRow] = deque() | |
| pending: deque[tuple[ContinuousRequest, float]] = deque() | |
| seen_ids: set[str] = set() | |
| input_done = False | |
| started = time.perf_counter() | |
| capacity = min(engine.max_batch_rows, engine.max_batch_tokens // | |
| int(engine.runtime.config.canvas_length)) | |
| def accept(item: Any, created: float) -> dict[str, Any] | None: | |
| nonlocal input_done | |
| if item is None: | |
| input_done = True | |
| elif isinstance(item, dict): | |
| return item | |
| elif item.request_id in seen_ids: | |
| return {"event": "request_error", "request_id": item.request_id, | |
| "error": "Duplicate request_id."} | |
| else: | |
| seen_ids.add(item.request_id) | |
| pending.append((item, created)) | |
| while not input_done or pending or prefilling or active: | |
| while not input_done and len(pending) < max_queue_size: | |
| try: | |
| error = accept(*incoming.get_nowait()) | |
| if error is not None: | |
| yield error | |
| except queue.Empty: | |
| break | |
| prefill_budget = engine.prefill_chunk | |
| had_active = bool(active) | |
| while prefill_budget > 0: | |
| if pending and len(active) + len(prefilling) < engine.max_batch_rows: | |
| request, created = pending.popleft() | |
| try: | |
| # A short newcomer gets one chunk promptly; incomplete | |
| # prompts then rotate in the bounded prefill cohort. | |
| prefilling.appendleft(engine.begin_prefill(request, created)) | |
| except Exception as error: | |
| yield {"event": "request_error", "request_id": request.request_id, | |
| "error": str(error)} | |
| continue | |
| if not prefilling: | |
| break | |
| work = prefilling.popleft() | |
| tokens = min(engine.prefill_chunk, len(work.prompt_ids) - work.offset) | |
| if tokens > prefill_budget: | |
| # Keep original chunk boundaries and singleton GEMM shapes. | |
| prefilling.appendleft(work) | |
| break | |
| prefill_budget -= tokens | |
| try: | |
| row = engine.prefill_step(work) | |
| if row is None: | |
| prefilling.append(work) | |
| else: | |
| active.append(row) | |
| except Exception as error: | |
| yield {"event": "request_error", "request_id": work.request.request_id, | |
| "error": str(error)} | |
| if not had_active and active: | |
| # Deliver the first request's initial denoise promptly. Once | |
| # decoding, use the remaining token budget to fill short rows. | |
| break | |
| if active: | |
| engine.peak_active_requests = max(getattr(engine, "peak_active_requests", 0), | |
| len(active) + len(prefilling)) | |
| selected = [active.popleft() for _ in range(min(capacity, len(active)))] | |
| try: | |
| finished = engine.step(selected) | |
| finished_ids = {id(row) for row in finished} | |
| for row in selected: | |
| if id(row) in finished_ids: | |
| yield engine.result(row) | |
| else: | |
| active.append(row) | |
| except Exception as error: | |
| for row in selected: | |
| yield {"event": "generation_result", "request_id": row.request.request_id, | |
| "status": "failed", "error": str(error)} | |
| if not continue_on_error: | |
| raise | |
| elif not pending and not prefilling and not input_done: | |
| # Wake immediately for new input instead of polling every 10 ms. | |
| error = accept(*incoming.get()) | |
| if error is not None: | |
| yield error | |
| tail_started = time.perf_counter() | |
| mx.synchronize() | |
| engine.denoise_seconds += time.perf_counter() - tail_started | |
| elapsed = time.perf_counter() - started | |
| yield {"event": "inference_summary", "elapsed_seconds": elapsed, | |
| "active_row_denoises": engine.active_row_steps, | |
| "heavy_forward_count": engine.forward_count, | |
| "denoise_steps_per_second": engine.active_row_steps / max(elapsed, 1e-9), | |
| "denoise_scheduler_seconds": engine.denoise_seconds, | |
| "scheduler_denoises_per_second": engine.active_row_steps / | |
| max(engine.denoise_seconds, 1e-9), | |
| "pipeline_depth": engine.pipeline_depth, | |
| "max_inflight_denoises": engine.max_inflight, | |
| "peak_active_requests": getattr(engine, "peak_active_requests", 0), | |
| "mlx_peak_memory_bytes": int(mx.get_peak_memory())} | |
| def _run_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any], | |
| max_queue_size: int) -> int: | |
| had_error = False | |
| for record in iter_scheduler(engine, incoming, max_queue_size): | |
| had_error |= record["event"] == "request_error" or record.get("status") == "failed" | |
| engine.emit(record) | |
| return int(had_error) | |
| def _parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser(description="Modilify Mk2 native MLX text generation") | |
| parser.add_argument("--model", "--checkpoint", dest="checkpoint", | |
| default=(str(Path(__file__).resolve().parent) | |
| if (Path(__file__).resolve().parent / "export_manifest.json").is_file() | |
| else "Modilify/Modilify-Mk2-preview-mlx"), | |
| help="Local model directory or Hugging Face Hub model ID.") | |
| parser.add_argument("--prompt", default="Why is the sky blue?") | |
| parser.add_argument("--requests-jsonl", metavar="PATH|-", | |
| help="Read JSONL requests continuously; '-' reads stdin.") | |
| parser.add_argument("--stream", action="store_true", | |
| help="Stream generated text to stdout for a single --prompt request.") | |
| parser.add_argument("--think", type=parse_bool, default=True) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--max-new-tokens", type=int, default=8192) | |
| parser.add_argument("--max-denoising-steps", type=int) | |
| parser.add_argument("--canvas-length", type=int) | |
| parser.add_argument("--repetition-penalty", type=float, default=1.0) | |
| parser.add_argument("--batch-size", type=int, default=4) | |
| parser.add_argument("--pipeline-depth", type=int, default=2, | |
| help="Bound in-flight singleton denoises; 1 minimizes activation memory.") | |
| parser.add_argument("--max-batch-tokens", type=int, default=1024) | |
| parser.add_argument("--prefill-chunk-size", type=int, default=256) | |
| parser.add_argument("--vocab-chunk-size", type=int, default=4096) | |
| parser.add_argument("--prefix-cache-mib", type=int, default=512) | |
| parser.add_argument("--max-queue-size", type=int, default=128) | |
| parser.add_argument("--commit-failure-budget", type=float, default=None, | |
| help="Override commit failure budget (defaults to checkpoint config).") | |
| parser.add_argument("--commit-top-k", type=int, default=None, | |
| help="Override commit top_k bound (defaults to checkpoint config).") | |
| parser.add_argument("--commit-min-p", type=float, default=None, | |
| help="Override commit min_p bound (defaults to checkpoint config).") | |
| parser.add_argument("--commit-target-confidence", type=float, default=None, | |
| help="Override commit target confidence (defaults to checkpoint config).") | |
| return parser | |
| def _validate_args(parser: argparse.ArgumentParser, args: Any) -> None: | |
| if args.stream and args.requests_jsonl is not None: | |
| parser.error("--stream can only be used with a single --prompt request.") | |
| positive = ("max_new_tokens", "batch_size", "max_batch_tokens", | |
| "prefill_chunk_size", "vocab_chunk_size", "max_queue_size", "pipeline_depth") | |
| for name in positive: | |
| if getattr(args, name) <= 0: | |
| parser.error(f"--{name.replace('_', '-')} must be positive.") | |
| if args.max_denoising_steps is not None and args.max_denoising_steps <= 0: | |
| parser.error("--max-denoising-steps must be positive.") | |
| if args.prefix_cache_mib < 0: | |
| parser.error("--prefix-cache-mib must be nonnegative.") | |
| if not math.isfinite(args.repetition_penalty) or args.repetition_penalty <= 0: | |
| parser.error("--repetition-penalty must be finite and positive.") | |
| if args.commit_failure_budget is not None and args.commit_failure_budget <= 0: | |
| parser.error("--commit-failure-budget must be positive.") | |
| if args.commit_top_k is not None and args.commit_top_k <= 0: | |
| parser.error("--commit-top-k must be positive.") | |
| if args.commit_min_p is not None and not (0.0 < args.commit_min_p < 1.0): | |
| parser.error("--commit-min-p must be in (0, 1).") | |
| if args.commit_target_confidence is not None and not (0.0 < args.commit_target_confidence < 1.0): | |
| parser.error("--commit-target-confidence must be in (0, 1).") | |
| def _producer(path: str, output: queue.Queue[Any], args: Any) -> None: | |
| stream = None | |
| try: | |
| stream = sys.stdin if path == "-" else open(path, encoding="utf-8") | |
| for line_number, line in enumerate(stream, 1): | |
| if not line.strip(): | |
| continue | |
| try: | |
| record = json.loads(line) | |
| request = parse_continuous_request( | |
| record, default_max_new_tokens=args.max_new_tokens, | |
| default_max_denoising_steps=args.max_denoising_steps, | |
| default_seed=args.seed, default_think=args.think, | |
| ) | |
| output.put((request, time.perf_counter())) | |
| except Exception as error: | |
| output.put(({"event": "request_error", "line": line_number, | |
| "error": str(error)}, time.perf_counter())) | |
| except Exception as error: | |
| output.put(({"event": "request_error", "error": str(error)}, | |
| time.perf_counter())) | |
| finally: | |
| if stream is not None and stream is not sys.stdin: | |
| stream.close() | |
| output.put((None, time.perf_counter())) | |
| def main(argv: list[str] | None = None) -> int: | |
| parser = _parser() | |
| args = parser.parse_args(argv) | |
| _validate_args(parser, args) | |
| runtime = load_runtime(args.checkpoint, canvas_length=args.canvas_length, | |
| max_new_tokens=args.max_new_tokens, | |
| max_denoising_steps=args.max_denoising_steps, | |
| repetition_penalty=args.repetition_penalty, | |
| commit_failure_budget=args.commit_failure_budget, | |
| commit_top_k=args.commit_top_k, | |
| commit_min_p=args.commit_min_p, | |
| commit_target_confidence=args.commit_target_confidence) | |
| if args.max_denoising_steps is None: | |
| args.max_denoising_steps = runtime.generation.max_denoising_steps | |
| streamed_token_ids: dict[str, list[int]] = {} | |
| streamed_text: dict[str, str] = {} | |
| def emit(record: dict[str, Any]) -> None: | |
| if args.stream: | |
| event = record.get("event") | |
| request_id = str(record.get("request_id", "single")) | |
| if event == "token": | |
| ids = streamed_token_ids.setdefault(request_id, []) | |
| ids.extend(int(token_id) for token_id in record["token_ids"]) | |
| current = runtime.tokenizer.decode(ids, skip_special_tokens=True) | |
| previous = streamed_text.get(request_id, "") | |
| if current.startswith(previous): | |
| delta = current[len(previous):] | |
| else: | |
| # Keep output append-only if a tokenizer revises a prior decode. | |
| common = 0 | |
| for old_char, new_char in zip(previous, current): | |
| if old_char != new_char: | |
| break | |
| common += 1 | |
| delta = current[common:] | |
| if delta: | |
| sys.stdout.write(delta) | |
| sys.stdout.flush() | |
| streamed_text[request_id] = current | |
| elif event == "generation_result": | |
| if record.get("status") == "failed": | |
| sys.stderr.write(f"\nInference failed: {record.get('error', 'unknown error')}\n") | |
| sys.stderr.flush() | |
| else: | |
| sys.stdout.write("\n") | |
| sys.stdout.flush() | |
| elif event == "request_error": | |
| sys.stderr.write(f"\nInference error: {record.get('error', 'unknown error')}\n") | |
| sys.stderr.flush() | |
| return | |
| sys.stdout.write(json.dumps(record, ensure_ascii=False) + "\n") | |
| sys.stdout.flush() | |
| emit({"event": "runtime_loaded", "checkpoint": str(runtime.checkpoint), | |
| "step": runtime.step, "restored_trainable_tensors": runtime.tensor_count, | |
| "backend": "mlx", "canvas_length": runtime.config.canvas_length, | |
| "pipeline_depth": args.pipeline_depth, | |
| "execution": "layer_interleaved_singleton", | |
| "commit_failure_budget": runtime.config.commit_failure_budget, | |
| "commit_target_confidence": getattr(runtime.config, "commit_target_confidence", None), | |
| "commit_top_k": getattr(runtime.config, "commit_top_k", None), | |
| "commit_min_p": getattr(runtime.config, "commit_min_p", None)}) | |
| engine = MLXContinuousEngine( | |
| runtime, prefix_mib=args.prefix_cache_mib, | |
| prefill_chunk=args.prefill_chunk_size, | |
| vocab_chunk=args.vocab_chunk_size, | |
| max_batch_rows=args.batch_size, | |
| max_batch_tokens=args.max_batch_tokens, emit=emit, pipeline_depth=args.pipeline_depth, | |
| ) | |
| incoming: queue.Queue[Any] = queue.Queue(maxsize=max(2, args.max_queue_size)) | |
| if args.requests_jsonl is None: | |
| request = ContinuousRequest("single", [{"role": "user", "content": args.prompt}], | |
| args.max_new_tokens, args.max_denoising_steps, | |
| args.seed, args.think, args.prompt) | |
| incoming.put((request, time.perf_counter())) | |
| incoming.put((None, time.perf_counter())) | |
| else: | |
| threading.Thread(target=_producer, args=(args.requests_jsonl, incoming, args), | |
| daemon=True).start() | |
| return _run_scheduler(engine, incoming, args.max_queue_size) | |
| def load_model( | |
| model: str = "Modilify/Modilify-Mk2-preview-mlx", *, | |
| canvas_length: int | None = None, | |
| ) -> MLXRuntime: | |
| """Load the complete local model or its Hugging Face snapshot once.""" | |
| return load_runtime(model, canvas_length=canvas_length, | |
| max_new_tokens=256, max_denoising_steps=None, | |
| repetition_penalty=1.0) | |
| def generate( | |
| runtime: MLXRuntime, prompt: str | None = None, *, | |
| messages: list[dict[str, Any]] | None = None, | |
| max_new_tokens: int = 256, max_denoising_steps: int | None = None, | |
| seed: int = 42, think: bool = True, | |
| ) -> dict[str, Any]: | |
| """Generate one independent response, returning text, tokens, and metrics.""" | |
| record = {"request_id": "single", "max_new_tokens": max_new_tokens, | |
| "max_denoising_steps": max_denoising_steps, "seed": seed, "think": think} | |
| if prompt is not None: | |
| record["prompt"] = prompt | |
| if messages is not None: | |
| record["messages"] = messages | |
| request = parse_continuous_request( | |
| record, default_max_new_tokens=256, | |
| default_max_denoising_steps=runtime.generation.max_denoising_steps, | |
| default_seed=42, default_think=True, | |
| ) | |
| engine = MLXContinuousEngine( | |
| runtime, prefix_mib=0, prefill_chunk=256, vocab_chunk=4096, | |
| max_batch_rows=1, max_batch_tokens=int(runtime.config.canvas_length), | |
| emit=lambda event: None, pipeline_depth=1, token_events=False, | |
| ) | |
| row = engine.admit(request, time.perf_counter()) | |
| while not engine.step([row]): | |
| pass | |
| return engine.result(row) | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |