"""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())