Repository navigation
Multimodality support - #2301
OnePunchMonk wants to merge 8 commits into
Conversation
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.
GPU test results (Modal, A10G)Ran this PR's branch ( Test suite
Total: 1193 passed, 0 failed across the vision, LoRA, Adapter, AdapterV2, base model, chat, and generate suites, confirming the Backward-compatibility checkNone of the existing suites explicitly call 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: |
|
@bhimrazy could you provide a first review of this when you have some time to spare? |
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
|
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
How I tested it Unit tests (
End-to-end smoke run against the real Regression run on CPU (Python 3.12):
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. |
|
Ran a full e2e pass against the current head ( Unit tests
Regression tests
Manual e2e script (tiny random-init multimodal model, no downloads needed):
No regressions found, everything passes cleanly on this rebase. |
closes #2173
Supersedes #2232 (which had gone stale against
mainand was hard to review after 4 months / 14 commits). This is the same feature rebased onto currentmainas a single commit, plus a fix for a bug found while chasing down a CI failure on the old PR: the LoRA/Adapter/AdapterV2GPTsubclasses construct their own state instead of callingGPT.__init__, so they never received the newvision_encoder/mm_projectorattributes — passingpixel_valuesthrough those model types raisedAttributeError.Summary
litgpt/vision.py:VisionEncoder(HF backbone or conv-patch fallback),MultiModalProjector(linear/mlp2x),merge_input_embeds,ImagePreprocessor.Configgains optionalvision_*fields and anis_multimodalproperty.GPT.forwardacceptspixel_valuesand merges projected image-patch embeddings into the token embedding sequence at<image>placeholder positions. Same wiring added to LoRA/Adapter/AdapterV2GPTsubclasses.litgpt.generate.base/litgpt.chat.basethreadpixel_valuesthrough generation;LLM.generate()accepts animagepath/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) coversVisionEncoder,MultiModalProjector,merge_input_embeds,ImagePreprocessor, and GPT integration.tests/test_lora.py,tests/test_adapter.py,tests/test_adapter_v2.py,tests/test_model.pylocally on CPU against latestmain— all pass.ruff checkpasses on all changed files.AI Usage Disclaimer