mlx-llada2-uni / llada2 /generate.py
treadon's picture
Upload llada2/generate.py with huggingface_hub
04482b3 verified
Raw History Blame Contribute Delete
6.5 kB
"""Block-diffusion (mask token prediction) text generation for LLaDA2.
Simplified port of generate_bd from the HF remote code:
- Fill a template with MASK tokens beyond the prompt
- Process block-by-block (left to right)
- Per block: N denoising steps; each step predicts all masked positions, accepts
the most-confident ones until the block is fully filled
- Block-diagonal "causal" mask over blocks: within a block we attend bidirectionally;
future blocks are not visible to the current block
"""
from __future__ import annotations
import math
from typing import Optional
import mlx.core as mx
def _build_block_mask(num_blocks: int, block_length: int) -> mx.array:
"""Lower-triangular mask over blocks, expanded to token-level.
Returns shape [1, 1, total, total] with 0 for allowed, -inf for masked.
MLX SDPA expects additive masks with -inf for disallowed positions.
"""
total = num_blocks * block_length
# block index per position
pos = mx.arange(total)
block_id = pos // block_length
# Allowed when key_block <= query_block
allowed = block_id[:, None] >= block_id[None, :] # [total, total]
# Convert to additive float mask
mask = mx.where(
allowed,
mx.array(0.0, dtype=mx.float32),
mx.array(-float("inf"), dtype=mx.float32),
)
return mask[None, None, :, :]
def _sample_greedy(logits: mx.array) -> mx.array:
return mx.argmax(logits, axis=-1)
def _softmax_probs(logits: mx.array) -> mx.array:
return mx.softmax(logits.astype(mx.float32), axis=-1)
def generate_text(
model,
prompt_ids: mx.array,
gen_length: int = 128,
block_length: int = 32,
steps_per_block: int = 16,
temperature: float = 0.0,
threshold: float = 0.95,
mask_token_id: int = 156895,
eos_token_id: int = 156892,
verbose: bool = True,
) -> mx.array:
"""Generate text tokens via block diffusion. prompt_ids: [1, L]."""
assert prompt_ids.ndim == 2 and prompt_ids.shape[0] == 1
prompt_len = prompt_ids.shape[1]
num_blocks = (prompt_len + gen_length + block_length - 1) // block_length
total_length = num_blocks * block_length
attn_mask = _build_block_mask(num_blocks, block_length).astype(mx.float32)
# Template: prompt || MASK...
x = mx.full((1, total_length), mask_token_id, dtype=mx.int32)
x[:, :prompt_len] = prompt_ids.astype(mx.int32)
prefill_blocks = prompt_len // block_length
# Per-step "how many tokens to transfer this step"
base = block_length // steps_per_block
remainder = block_length % steps_per_block
transfer_schedule = [base + (1 if i < remainder else 0) for i in range(steps_per_block)]
for block_idx in range(prefill_blocks, num_blocks):
current_end = (block_idx + 1) * block_length
block_start = block_idx * block_length
cur_x = x[:, :current_end]
cur_mask = attn_mask[:, :, :current_end, :current_end]
if verbose:
print(f"[gen] block {block_idx - prefill_blocks + 1}/{num_blocks - prefill_blocks}")
for step in range(steps_per_block):
# Positions still masked within current block
block_slice = cur_x[:, block_start:current_end]
is_masked = block_slice == mask_token_id
if not bool(is_masked.any()):
break
logits = model(cur_x, attn_mask=cur_mask) # [1, current_end, V]
block_logits = logits[:, block_start:current_end, :] # [1, block_length, V]
if temperature == 0.0:
predicted = _sample_greedy(block_logits) # [1, block_length]
probs = _softmax_probs(block_logits)
confidence = mx.take_along_axis(probs, predicted[..., None], axis=-1).squeeze(-1)
else:
# temperature sampling (not used for determinism during smoke test)
scaled = block_logits / temperature
probs = _softmax_probs(scaled)
predicted = _sample_greedy(scaled)
confidence = mx.take_along_axis(probs, predicted[..., None], axis=-1).squeeze(-1)
confidence = mx.where(is_masked, confidence, mx.array(-float("inf"), dtype=confidence.dtype))
# Decide how many to transfer this step
num_to_transfer = transfer_schedule[step]
# Always include "high-confidence" tokens above threshold
high_conf = confidence > threshold
n_high = int(high_conf.sum().item())
# transfer_index: positions to accept this step
if n_high >= num_to_transfer:
transfer_index = high_conf
else:
# pick top num_to_transfer by confidence (clamped to number of still-masked)
n_masked = int(is_masked.sum().item())
k = min(num_to_transfer, n_masked)
# top-k across block axis
conf_flat = confidence[0] # [block_length]
# argpartition -> top k (largest). negate for descending.
idx = mx.argpartition(-conf_flat, k - 1)[:k]
transfer_index = mx.zeros(conf_flat.shape, dtype=mx.bool_)
transfer_index = mx.put_along_axis(
transfer_index, idx, mx.ones((k,), dtype=mx.bool_), axis=-1
)
transfer_index = transfer_index[None, :]
# Apply transfers
new_block = mx.where(transfer_index, predicted, block_slice)
# Write back into cur_x and x
cur_x = mx.concatenate([cur_x[:, :block_start], new_block], axis=1)
# Early stop on EOS once anchored in a filled span
if bool((new_block == eos_token_id).any()):
filled = new_block != mask_token_id
if bool(filled.all()):
break
# Commit block to x
x = mx.concatenate([cur_x, x[:, current_end:]], axis=1)
# Optional global early stop on EOS
gen_so_far = x[0, prompt_len:current_end]
if bool((gen_so_far == eos_token_id).any()):
break
# Trim at first EOS after prompt
gen = x[0, prompt_len:]
eos_positions = mx.argmax((gen == eos_token_id).astype(mx.int32), axis=0)
has_eos = bool((gen == eos_token_id).any())
if has_eos:
cutoff = int(eos_positions.item())
return x[:, : prompt_len + cutoff + 1]
return x[:, : prompt_len + gen_length]