Skip to content

Multimodality support - #2301

Open
OnePunchMonk wants to merge 8 commits into
Lightning-AI:mainfrom
OnePunchMonk:multimodality-support-v2
Open

OnePunchMonk wants to merge 8 commits into
Lightning-AI:mainfrom
OnePunchMonk:multimodality-support-v2

Conversation

@OnePunchMonk

@OnePunchMonk OnePunchMonk commented Aug 17, 2026 •

Copy link
Copy Markdown
Contributor

closes #2173

Supersedes #2232 (which had gone stale against main and was hard to review after 4 months / 14 commits). This is the same feature rebased onto current main as a single commit, plus a fix for a bug found while chasing down a CI failure on the old PR: the LoRA/Adapter/AdapterV2 GPT subclasses construct their own state instead of calling GPT.__init__, so they never received the new vision_encoder/mm_projector attributes — passing pixel_values through those model types raised AttributeError.

Summary

  • Adds litgpt/vision.py: VisionEncoder (HF backbone or conv-patch fallback), MultiModalProjector (linear/mlp2x), merge_input_embeds, ImagePreprocessor.
  • Config gains optional vision_* fields and an is_multimodal property.
  • GPT.forward accepts pixel_values and merges projected image-patch embeddings into the token embedding sequence at <image> placeholder positions. Same wiring added to LoRA/Adapter/AdapterV2 GPT subclasses.
  • litgpt.generate.base / litgpt.chat.base thread pixel_values through generation; LLM.generate() accepts an image path/PIL input.
  • convert_hf_checkpoint: only load Gemma3 vision tower / mm-projector weights when the target config explicitly declares a matching vision architecture, to avoid corrupting the state_dict when shapes don't match Gemma3's SigLIP tower.

Test plan

  • tests/test_vision.py (new, 23 tests) covers VisionEncoder, MultiModalProjector, merge_input_embeds, ImagePreprocessor, and GPT integration.
  • Ran full tests/test_lora.py, tests/test_adapter.py, tests/test_adapter_v2.py, tests/test_model.py locally on CPU against latest main — all pass.
  • ruff check passes on all changed files.

AI Usage Disclaimer

  • AI assistance (Claude Code) was used for this change.

Adds a vision encoder + multimodal projector pipeline for VLMs:

- litgpt/vision.py: VisionEncoder (HF backbone or conv fallback),
  MultiModalProjector (linear/mlp2x), merge_input_embeds, and
  ImagePreprocessor.
- Config gains optional vision_* fields and an is_multimodal property.
- GPT.forward accepts pixel_values and merges projected image patch
  embeddings into the token embedding sequence at <image> placeholder
  positions. Same wiring is added to the LoRA/Adapter/AdapterV2 GPT
  subclasses, which construct their own state rather than calling
  GPT.__init__.
- litgpt.generate.base and litgpt.chat.base thread pixel_values
  through generation; LLM.generate() accepts an `image` path/PIL input.
- convert_hf_checkpoint: load Gemma3 vision tower / mm-projector
  weights only when the target config explicitly declares a matching
  vision architecture, to avoid corrupting the state_dict when it
  doesn't match Gemma3's SigLIP tower shapes.
@OnePunchMonk

OnePunchMonk commented Aug 23, 2026 •

Copy link
Copy Markdown
Contributor Author

GPU test results (Modal, A10G)

Ran this PR's branch (multimodality-support-v2, 0f4e468) end-to-end on a real GPU via Modal, installing with pip install .[extra,test,compiler] on Python 3.11 / an NVIDIA A10.

Test suite

File Result
tests/test_vision.py 23 passed
tests/test_lora.py 392 passed, 3 skipped
tests/test_adapter.py 15 passed, 6 xfailed
tests/test_adapter_v2.py 226 passed
tests/test_model.py 515 passed, 7 skipped, 62 xfailed, 22 xpassed
tests/test_chat.py 13 passed
tests/generate/test_main.py 9 passed, 1 xfailed

Total: 1193 passed, 0 failed across the vision, LoRA, Adapter, AdapterV2, base model, chat, and generate suites, confirming the GPT/LoRA GPT/Adapter GPT/AdapterV2 GPT vision wiring (mentioned in the PR description) doesn't regress any of those paths on GPU.

Backward-compatibility check

None of the existing suites explicitly call GPT.forward both with and without the new pixel_values kwarg on a non-multimodal config, so I added a standalone check:

config = Config.from_name("pythia-14m")  # no vision_* fields
model = GPT(config)
...
out_no_kwarg = model(idx)                    # old call signature
out_none_kwarg = model(idx, pixel_values=None)  # new call signature
assert torch.equal(out_no_kwarg, out_none_kwarg)

Result: BACKWARD_COMPAT_OK — identical output whether pixel_values is omitted or passed as None, so existing non-multimodal checkpoints are unaffected by this change.

@OnePunchMonk

Copy link
Copy Markdown
Contributor Author

@bhimrazy could you provide a first review of this when you have some time to spare?

OnePunchMonk and others added 6 commits October 2, 2026 14:01
no_grad() was wrapping the entire VisionEncoder.forward, including the
trainable Conv2d fallback path used when no HF encoder is loaded. Scope
no_grad to only the frozen HF-encoder branch.
Remove trailing comma that forced ruff's magic-trailing-comma rule to
explode a call that fits on one line, matching the identical call in
chat/base.py.
- Expand <image> placeholders to num_patches tokens in LLM.generate and chat
- Use the vision tower of CLIP/SigLIP checkpoints loaded via AutoModel
- Require vision_start_token_id when vision_feature_dim is set
- Drop Gemma3 vision weight loading from convert_hf_checkpoint (shapes did not match)
…e vision init

- Build the HF vision tower from its config so models can be created on the meta device;
  load pretrained weights on real-device init or when a checkpoint lacks encoder weights
- Add ImagePreprocessor.from_config to take size and mean/std from the HF image processor;
  rename IMAGENET_* constants to CLIP_*
- Move the vision encoder/projector setup into GPT._init_vision, used by LoRA/Adapter/AdapterV2
@OnePunchMonk

Copy link
Copy Markdown
Contributor Author

Pushed two more commits (24934db, f4703e9) after going through the diff again. A few things in the earlier version didn't actually work end to end, so here's what changed and how I checked it.

What was broken and is now fixed

  1. VisionEncoder loaded the HF model with AutoModel.from_pretrained, which for CLIP/SigLIP returns the full dual-tower model. Its forward needs input_ids, so passing only pixel_values failed with ValueError: You have to specify input_ids. It now keeps only the vision_model.
  2. Nothing inserted the <image> placeholder tokens, so LLM.generate(image=...) and litgpt chat --image always hit the placeholder count check in merge_input_embeds. Added expand_image_tokens(): a single placeholder in the prompt is expanded to num_patches copies, and if there's none the block goes right after BOS.
  3. The Gemma3 vision weight loading in convert_hf_checkpoint.py mapped keys onto modules that don't match Gemma3's SigLIP tower or its projector, so I reverted that file to main. Gemma3 stays text-only here and vision weight conversion can be a separate PR.
  4. Config now raises if vision_feature_dim is set without vision_start_token_id.
  5. The HF vision tower is built from its config instead of downloading weights in __init__, so the model can be created on the meta device. Pretrained weights are loaded on real-device init, or by a load_state_dict pre-hook when the checkpoint has no encoder weights (e.g. a lit checkpoint converted from the text weights), so strict loading still works.
  6. ImagePreprocessor.from_config(config) takes image size and mean/std from the HF image processor (SigLIP uses 0.5, CLIP its own values). The old IMAGENET_* constants were actually CLIP's values, so they're renamed to CLIP_*.
  7. The vision setup that was copy-pasted into the four GPT classes is now one GPT._init_vision() that LoRA/Adapter/AdapterV2 call.

How I tested it

Unit tests (tests/test_vision.py, now 38 tests, no network needed since HF loading is mocked with a tiny random CLIP):

  • expand_image_tokens: insert after BOS, insert with no BOS, expand a single placeholder in place, already-expanded prompt unchanged, wrong count raises, and the result runs through GPT.forward with pixel_values
  • HF encoder: uses the vision tower, weights match the pretrained ones, encoder is frozen
  • meta-device init doesn't call from_pretrained
  • strict load_state_dict of a checkpoint with no vision_encoder._encoder.* keys fills them from HF. I checked this test fails with "Missing key(s)" when the hook is disabled.
  • ImagePreprocessor.from_config with and without an HF processor
  • LoRA/Adapter/AdapterV2 GPT get the vision components from a multimodal config

End-to-end smoke run against the real hf-internal-testing/tiny-random-CLIPModel and tiny-random-SiglipModel: build GPT on meta, strict-load a checkpoint without encoder weights, preprocess a PIL image with from_config, expand tokens, forward. Both give (1, 228, 512) logits with all finite values, and the preprocessor picks size 30 with mean 0.481 (CLIP) / 0.5 (SigLIP).

Regression run on CPU (Python 3.12):

pytest tests/test_vision.py tests/test_lora.py tests/test_adapter.py tests/test_adapter_v2.py \
  tests/test_model.py tests/test_chat.py tests/test_api.py tests/test_config.py \
  tests/convert/test_hf_checkpoint.py tests/generate/test_main.py
1032 passed, 524 skipped, 2 xfailed

ruff check and ruff format --check are clean. All 524 skips are _RunIf(min_cuda_gpus=...) tests, since this was a CPU-only machine. The earlier A10G GPU run is in my comment above, but it predates these two commits.

Things I left out on purpose: no Gemma3/LLaVA weight conversion and no multi-image prompts. I'd rather do those as follow-ups once the base wiring here looks right to you.

@OnePunchMonk OnePunchMonk reopened this Oct 5, 2026
@OnePunchMonk

Copy link
Copy Markdown
Contributor Author

Ran a full e2e pass against the current head (889a50d):

Unit tests

  • tests/test_vision.py: 38/38 passed

Regression tests

  • test_lora.py, test_adapter.py, test_adapter_v2.py, test_model.py: 729 passed, 0 failed (skips are CUDA/bf16 gated, expected on CPU)

Manual e2e script (tiny random-init multimodal model, no downloads needed):

  • GPT.forward(idx, pixel_values=...) merges image patches correctly and produces the right output shape
  • Same call through LoRAGPT, AdapterGPT, AdapterV2GPT subclasses works, confirming the fix for the AttributeError these subclasses used to hit (they build their own state instead of calling GPT.__init__)
  • litgpt.generate.base.generate(..., pixel_values=...) runs full autoregressive generation with KV cache and image conditioning
  • ImagePreprocessor on a real PIL.Image produces correctly resized/normalized output

No regressions found, everything passes cleanly on this rebase.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Feature: Support for Multimodality

1 participant