Any-to-Any
MLX
diffusion-lm
mixture-of-experts
multimodal
text-to-image
image-understanding
apple-silicon
llada
Instructions to use treadon/mlx-llada2-uni with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use treadon/mlx-llada2-uni with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir mlx-llada2-uni treadon/mlx-llada2-uni
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download llada2/generate.py from treadon/mlx-llada2-uni: direct link, hf CLI and curl.
- Browser
- Download file 6.5 kB
-
https://huggingface.co/treadon/mlx-llada2-uni/resolve/main/llada2/generate.py
- Command line
-
hf download hf://treadon/mlx-llada2-uni/llada2/generate.py
-
curl -L -o generate.py https://huggingface.co/treadon/mlx-llada2-uni/resolve/main/llada2/generate.py
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] | |