Modilify-Mk2-preview-mlx / inference.py
ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
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",
}
)
@dataclass(frozen=True)
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
@dataclass
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
@dataclass
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
@dataclass
class PrefillRow:
request: ContinuousRequest
prompt_ids: tuple[int, ...]
cache: list[Any]
offset: int
reused: int
created: float
seconds: float = 0.0
@dataclass
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())