ydy9038074's picture
Publish Modilify Mk1 MLX runtime
a066584 verified
Raw
History Blame Contribute Delete
15.8 kB
# Copyright 2026 Modilify
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
"""Rolling canvas generation for native Modilify Mk1."""
from __future__ import annotations
from dataclasses import dataclass, replace
import math
import time
from typing import Any
import mlx.core as mx
from mlx_vlm.generate.common import wired_limit
from .commit_policy import fused_commit_failure_rate, select_commit_lengths
from .latent_deliberation import LatentDeliberationState
from .modeling import ModilifyMk1ForBlockDiffusion
@dataclass
class ModilifyMk1GenerationOutput:
sequences: list[int]
generated_ids: list[int]
generated_length: int
denoise_steps: int
jump_count: int
average_commit_len: float
stop_reason: str
tokens_per_forward: float
prefill_seconds: float = 0.0
generate_seconds: float = 0.0
first_denoise_seconds: float = 0.0
heavy_denoise_per_second: float = 0.0
steady_heavy_denoise_per_second: float = 0.0
tokens_per_second: float = 0.0
@dataclass
class _RollingState:
canvas: mx.array
confidence: mx.array
entropy: mx.array
age: mx.array
latent_state: LatentDeliberationState
history_hidden_state: mx.array | None
def _flatten_token_ids(*values: object) -> tuple[int, ...]:
token_ids: list[int] = []
for value in values:
if value is None:
continue
if isinstance(value, int):
token_ids.append(int(value))
elif isinstance(value, (list, tuple, set)):
token_ids.extend(int(token_id) for token_id in value if token_id is not None)
return tuple(dict.fromkeys(token_ids))
def _shift_prefix(tensor: mx.array, committed: int, fill_value: float | int) -> mx.array:
if committed <= 0:
return tensor
canvas = tensor.shape[1]
if committed >= canvas:
return mx.full(tensor.shape, fill_value, dtype=tensor.dtype)
kept = tensor[:, committed:]
fill_shape = (tensor.shape[0], committed, *tensor.shape[2:])
fill = mx.full(fill_shape, fill_value, dtype=tensor.dtype)
return mx.concatenate([kept, fill], axis=1)
def _shift_state(
state: _RollingState,
committed: int,
*,
vocab_size: int,
canvas_length: int,
unknown_entropy: float,
) -> _RollingState:
if committed <= 0:
return state
tail = mx.random.randint(0, vocab_size, (state.canvas.shape[0], committed))
canvas = mx.concatenate([state.canvas[:, committed:], tail], axis=1)
if canvas.shape[1] != canvas_length:
raise RuntimeError("Canvas shift produced an unexpected length.")
latent = state.latent_state
shifted_latent = LatentDeliberationState(
token_latents=_shift_prefix(latent.token_latents, committed, 0),
memory_slots=latent.memory_slots,
confidence=_shift_prefix(latent.confidence, committed, 0),
entropy=_shift_prefix(latent.entropy, committed, unknown_entropy),
age=_shift_prefix(latent.age, committed, 0),
token_changed=_shift_prefix(latent.token_changed, committed, 0),
confidence_delta=_shift_prefix(latent.confidence_delta, committed, 0),
entropy_delta=_shift_prefix(latent.entropy_delta, committed, 0),
ponder_steps=mx.zeros_like(latent.ponder_steps),
stagnation_steps=mx.zeros_like(latent.stagnation_steps),
)
history = state.history_hidden_state
if history is not None:
history = _shift_prefix(history, committed, 0)
return _RollingState(
canvas=canvas,
confidence=_shift_prefix(state.confidence, committed, 0),
entropy=_shift_prefix(state.entropy, committed, unknown_entropy),
age=_shift_prefix(state.age, committed, 0),
latent_state=shifted_latent,
history_hidden_state=history,
)
def generate(
model: ModilifyMk1ForBlockDiffusion,
input_ids: mx.array,
*,
max_new_tokens: int | None = None,
temperature: float | None = None,
attention_mask: mx.array | None = None,
pixel_values: mx.array | None = None,
mm_token_type_ids: mx.array | None = None,
max_denoising_steps: int | None = None,
seed: int | None = None,
profiler=None,
) -> ModilifyMk1GenerationOutput:
"""Generate one response with rolling block diffusion.
Batch size 1 only. Semantics follow ``generation_modilify_mk1.py``.
"""
if input_ids.ndim != 2 or input_ids.shape[0] != 1:
raise ValueError("Native Mk1 generation currently requires shape [1, sequence].")
if seed is not None:
mx.random.seed(int(seed))
config = model.config
canvas_length = int(config.canvas_length)
vocab_size = int(config.vocab_size)
dtype = model.model.decoder.embed_tokens.weight.dtype
max_new = int(max_new_tokens if max_new_tokens is not None else 256)
if max_new <= 0:
raise ValueError("`max_new_tokens` must be positive.")
denoise_temperature = (
float(config.denoise_temperature) if temperature is None else float(temperature)
)
max_iterations = max(1, max_new * int(config.max_ponder_steps))
unknown_entropy = math.log(vocab_size)
if attention_mask is None:
attention_mask = mx.ones(input_ids.shape, dtype=mx.bool_)
prefill_started = time.perf_counter()
cache = model.make_cache(max_size=int(input_ids.shape[1]) + max_new)
cache = model.prefill(
input_ids,
attention_mask=attention_mask,
cache=cache,
pixel_values=pixel_values,
mm_token_type_ids=mm_token_type_ids,
)
mx.eval([item for block in cache for item in getattr(block, "state", ())])
prefill_seconds = time.perf_counter() - prefill_started
compiled_decoder_step = None
latent = LatentDeliberationState.empty(
batch_size=1,
canvas_length=canvas_length,
latent_dim=config.latent_dim,
memory_slots=config.latent_memory_slots,
dtype=dtype,
)
state = _RollingState(
canvas=mx.random.randint(0, vocab_size, (1, canvas_length)),
confidence=mx.zeros((1, canvas_length), dtype=mx.float32),
entropy=mx.full((1, canvas_length), unknown_entropy, dtype=mx.float32),
age=mx.zeros((1, canvas_length), dtype=mx.int32),
latent_state=latent,
history_hidden_state=None,
)
stop_token_ids = _flatten_token_ids(config.turn_end_token_id, config.eos_token_id)
turn_end = int(config.turn_end_token_id)
generated: list[int] = []
denoise_steps = 0
jumps = 0
shifts = 0
stop_reason = "episode_watchdog"
prompt_ids = [int(token) for token in input_ids[0].tolist()]
generate_started = time.perf_counter()
first_denoise_seconds = 0.0
always_active = mx.array([True])
with wired_limit(model, None):
while len(generated) < max_new:
step_started = time.perf_counter()
remaining = mx.array([max_new - len(generated)], dtype=mx.int32)
output = model(
decoder_input_ids=state.canvas,
cache=cache,
previous_confidence=state.confidence,
previous_entropy=state.entropy,
token_age=state.age,
latent_state=state.latent_state,
history_hidden_state=state.history_hidden_state,
return_proposal_statistics=True,
denoise_temperature=denoise_temperature,
repetition_penalty=float(config.repetition_penalty),
compiled_decoder_step=None,
profiler=profiler,
)
denoise_steps += 1
proposal = output.proposal
proposal_confidence = output.proposal_confidence
token_entropy = output.token_entropy
greedy_proposal = output.greedy_proposal
greedy_confidence = output.greedy_confidence
next_canvas = proposal
next_confidence = proposal_confidence.astype(mx.float32)
next_entropy = token_entropy.astype(mx.float32)
next_latent = replace(
output.next_latent_state,
confidence=next_confidence,
entropy=next_entropy,
age=state.age + 1,
token_changed=(next_canvas != state.canvas).astype(mx.float32),
confidence_delta=next_confidence - state.confidence,
entropy_delta=next_entropy - state.entropy,
)
next_state = _RollingState(
canvas=next_canvas,
confidence=next_confidence,
entropy=next_entropy,
age=state.age + 1,
latent_state=next_latent,
history_hidden_state=output.heavy_hidden_state,
)
if profiler is not None:
commit_span = profiler.measure(
"commit",
proposal,
proposal_confidence,
token_entropy,
)
policy = select_commit_lengths(
sampled_token_ids=proposal,
normal_failure_rate=fused_commit_failure_rate(
proposal_confidence, token_entropy, vocab_size=vocab_size
),
previous_failure_rate=fused_commit_failure_rate(
state.confidence, state.entropy, vocab_size=vocab_size
),
greedy_token_ids=greedy_proposal,
jump_failure_rate=fused_commit_failure_rate(
greedy_confidence, token_entropy, vocab_size=vocab_size
),
ponder_steps=state.latent_state.ponder_steps,
stagnation_steps=state.latent_state.stagnation_steps,
active_rows=always_active,
remaining_lengths=remaining,
failure_budget=float(config.commit_failure_budget),
jump_failure_budget=float(config.jump_failure_budget),
stop_token_id=stop_token_ids,
max_ponder_steps=int(config.max_ponder_steps),
stagnation_threshold=int(config.jump_on_no_progress_after),
min_progress=float(config.min_trajectory_progress),
)
if profiler is not None:
commit_span.done(policy.commit_lengths, policy.jump_rows)
if profiler is not None:
span = profiler.measure(
"sync",
policy.commit_lengths,
policy.jump_rows,
policy.commit_token_ids,
)
mx.eval(
policy.commit_lengths,
policy.jump_rows,
policy.commit_token_ids,
)
commit_len = int(policy.commit_lengths[0].item())
jump = bool(policy.jump_rows[0].item())
if profiler is not None:
span.done()
next_state = replace(
next_state,
latent_state=replace(
next_state.latent_state,
ponder_steps=policy.ponder_steps,
stagnation_steps=policy.stagnation_steps,
),
)
if jump:
next_state = replace(next_state, canvas=policy.commit_token_ids)
jumps += 1
committed_ids: list[int] = []
if commit_len:
committed_ids = [
int(token)
for token in policy.commit_token_ids[0, :commit_len].tolist()
]
generated.extend(committed_ids)
committed_block = policy.commit_token_ids[:, :commit_len]
if profiler is not None:
span = profiler.measure("update_cache", committed_block)
cache = model.update_cache(committed_block, cache=cache)
mx.eval(
*[
block.keys
for block in cache
if getattr(block, "keys", None) is not None
],
*[
block.values
for block in cache
if getattr(block, "values", None) is not None
],
)
if profiler is not None:
span.done()
shifts += 1
state = _shift_state(
next_state,
commit_len,
vocab_size=vocab_size,
canvas_length=canvas_length,
unknown_entropy=unknown_entropy,
)
if profiler is not None:
profiler.finish_step()
if denoise_steps == 1:
first_denoise_seconds = time.perf_counter() - step_started
generate_started = time.perf_counter()
if committed_ids:
if turn_end in committed_ids:
stop_reason = "turn_end"
break
if any(
token in stop_token_ids and token != turn_end
for token in committed_ids
):
stop_reason = "eos"
break
if len(generated) >= max_new:
stop_reason = "max_new_tokens"
break
if (
max_denoising_steps is not None
and denoise_steps >= max_denoising_steps
):
stop_reason = "max_denoising_steps"
break
if denoise_steps >= max_iterations:
stop_reason = "episode_watchdog"
break
generate_seconds = time.perf_counter() - generate_started
if denoise_steps <= 1:
generate_seconds = first_denoise_seconds
steady_hdps = 0.0
overall_hdps = (1.0 / first_denoise_seconds) if first_denoise_seconds else 0.0
tokens_per_second = (
len(generated) / first_denoise_seconds if first_denoise_seconds else 0.0
)
else:
steady_hdps = (denoise_steps - 1) / max(generate_seconds, 1e-9)
overall_hdps = denoise_steps / max(
first_denoise_seconds + generate_seconds, 1e-9
)
tokens_per_second = len(generated) / max(
first_denoise_seconds + generate_seconds, 1e-9
)
average_commit = (len(generated) / shifts) if shifts else 0.0
tpf = (len(generated) / denoise_steps) if denoise_steps else 0.0
return ModilifyMk1GenerationOutput(
sequences=prompt_ids + generated,
generated_ids=generated,
generated_length=len(generated),
denoise_steps=denoise_steps,
jump_count=jumps,
average_commit_len=average_commit,
stop_reason=stop_reason,
tokens_per_forward=tpf,
prefill_seconds=prefill_seconds,
generate_seconds=first_denoise_seconds + generate_seconds,
first_denoise_seconds=first_denoise_seconds,
heavy_denoise_per_second=overall_hdps,
steady_heavy_denoise_per_second=steady_hdps,
tokens_per_second=tokens_per_second,
)