mlx-llada2-uni / image_understand.py
treadon's picture
Upload image_understand.py with huggingface_hub
669b8a8 verified
Raw History Blame Contribute Delete
6.05 kB
"""Image Understanding (VQA) — hybrid PyTorch + MLX.
- PyTorch image_tokenizer (ViT + VQVAE, 2.4 GB) encodes PIL image → VQ token IDs.
- MLX LLaDA2 backbone runs the block-diffusion text generation with the VQ-in-vocab
tokens spliced into the prompt.
The ViT/VQVAE loaded to PyTorch MPS is freed before MLX forward passes to stay
inside the ~64 GB unified memory budget.
"""
import argparse
import gc
import json
import os
import sys
import time
from pathlib import Path
import mlx.core as mx
from huggingface_hub import snapshot_download
from PIL import Image
from transformers import AutoTokenizer
# Resolve official repo path (sibling to this package)
REPO_ROOT = Path(__file__).resolve().parent.parent / "llada2-uni-repo"
sys.path.insert(0, str(REPO_ROOT))
# Stub flash_attn (not on Apple Silicon). The decoder's dispatch_attention_fn
# fallback handles attention via diffusers + SDPA.
import types as _types, importlib.machinery as _im
if "flash_attn" not in sys.modules:
_stub = _types.ModuleType("flash_attn")
_stub.__spec__ = _im.ModuleSpec(name="flash_attn", loader=None)
_stub.__version__ = "0.0.0-stub"
_stub.flash_attn_func = lambda *a, **k: (_ for _ in ()).throw(
RuntimeError("flash_attn unavailable"))
sys.modules["flash_attn"] = _stub
from llada2.model import LLaDA2Config, LLaDA2Model
from llada2.weights import load_weights_into_model
from llada2.generate import generate_text
def encode_image(image_path: str, model_dir: Path):
"""Return (token_ids, h, w) where tokens are VQ indices (no offset)."""
import torch
# Official encoder expects the dir layout of the HF snapshot.
from encoder.image_tokenizer import ImageTokenizer
from decoder.utils import generate_crop_size_list, var_center_crop
# Use CPU for image tokenizer — it's only 2.4 GB but MPS can OOM on
# concurrent allocations. CPU path works reliably and takes <30s.
use_mps = os.environ.get("LLADA2_ENCODER_DEVICE", "cpu") == "mps"
device = torch.device("mps" if use_mps and torch.backends.mps.is_available() else "cpu")
dtype = torch.bfloat16 if device.type == "mps" else torch.float32
print(f"[encode] loading ImageTokenizer on {device}…")
t0 = time.time()
tokenizer = ImageTokenizer(model_path=str(model_dir), device=str(device), dtype=dtype)
print(f"[encode] loaded in {time.time()-t0:.1f}s")
# Default crop target: 512x512 with 32-multiple aspect ratios (matches official script)
crop_sizes = generate_crop_size_list((512 // 32) ** 2, 32)
pil = var_center_crop(Image.open(image_path).convert("RGB"), crop_size_list=crop_sizes)
print(f"[encode] cropped image to {pil.size}")
info = tokenizer.encode_with_info(pil)
t, h, w = info["grid_thw"]
print(f"[encode] VQ grid: {t}x{h}x{w}, {info['num_tokens']} tokens")
# Free the PyTorch model before MLX loads
del tokenizer
gc.collect()
return info["token_ids"], h, w
def build_prompt(tokenizer, image_tokens: list[int], image_h: int, image_w: int,
question: str, offset: int) -> list[int]:
"""<|image|><h-token><w-token><boi>[image tokens][<|/image|>] [question]"""
soi = tokenizer("<|image|>").input_ids
eoi = tokenizer("<|/image|>").input_ids
boi = tokenizer("<boi>").input_ids
h_tok = tokenizer(f"<|reserved_token_{image_h}|>").input_ids
w_tok = tokenizer(f"<|reserved_token_{image_w}|>").input_ids
pfx = tokenizer(question).input_ids if question else []
img_vocab = [t + offset for t in image_tokens]
return soi + h_tok + w_tok + boi + img_vocab + eoi + pfx
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--image", required=True, type=str)
ap.add_argument("--question", default="Describe this image in detail.", type=str)
ap.add_argument("--gen-length", default=256, type=int)
ap.add_argument("--block-length", default=32, type=int)
ap.add_argument("--steps-per-block", default=16, type=int)
ap.add_argument("--threshold", default=0.95, type=float)
ap.add_argument("--repo-id", default="inclusionAI/LLaDA2.0-Uni", type=str)
args = ap.parse_args()
print("[mmu] fetching model files…", flush=True)
snap = Path(snapshot_download(
args.repo_id,
allow_patterns=[
"model-*.safetensors", "model.safetensors.index.json",
"config.json", "tokenizer*", "special_tokens_map.json",
"image_tokenizer/*",
],
))
print(f"[mmu] snap dir: {snap}", flush=True)
# ---------- Phase 1: encode image to VQ tokens in PyTorch ----------
image_tokens, h, w = encode_image(args.image, snap)
# ---------- Phase 2: run MLX backbone ----------
tokenizer = AutoTokenizer.from_pretrained(str(snap), trust_remote_code=True)
config = LLaDA2Config.from_hf(json.loads((snap / "config.json").read_text()))
model = LLaDA2Model(config)
print("[mmu] loading MLX backbone weights…")
t0 = time.time()
load_weights_into_model(model, snap, dtype=mx.bfloat16, verbose=False)
print(f"[mmu] backbone loaded in {time.time()-t0:.1f}s")
ids = build_prompt(tokenizer, image_tokens, h, w, args.question, config.image_token_offset)
prompt_ids = mx.array([ids], dtype=mx.int32)
print(f"[mmu] prompt token count: {len(ids)} (image tokens: {len(image_tokens)}, question: '{args.question}')")
t0 = time.time()
out = generate_text(
model, prompt_ids,
gen_length=args.gen_length,
block_length=args.block_length,
steps_per_block=args.steps_per_block,
temperature=0.0, threshold=args.threshold,
mask_token_id=config.mask_token_id, eos_token_id=config.eos_token_id,
verbose=True,
)
mx.eval(out)
dt = time.time() - t0
gen_ids = out[0, len(ids):].tolist()
text = tokenizer.decode(gen_ids, skip_special_tokens=True)
print(f"\n{'='*60}")
print(f"Q: {args.question}")
print(f"A: {text}")
print(f"{'='*60}")
print(f"(generated in {dt:.1f}s)")
if __name__ == "__main__":
main()