"""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]