From 2394b231c226a91901bbf1912d5aefa6c61f9eec Mon Sep 17 00:00:00 2001 From: Alison Shao <54658187+alisonshao@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:11:54 -0700 Subject: [PATCH] Fix Mistral3 retaining every vision-tower layer to read one (#39185) Co-authored-by: Xinyuan Tong Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> --- python/sglang/srt/models/llava.py | 18 ++- python/sglang/srt/models/mistral.py | 18 ++- .../models/test_mistral3_vision_feature.py | 109 ++++++++++++++++++ 3 files changed, 133 insertions(+), 12 deletions(-) create mode 100644 test/registered/unit/models/test_mistral3_vision_feature.py diff --git a/python/sglang/srt/models/llava.py b/python/sglang/srt/models/llava.py index c5f62f6b2..71d562869 100644 --- a/python/sglang/srt/models/llava.py +++ b/python/sglang/srt/models/llava.py @@ -804,16 +804,22 @@ class LlavaForConditionalGeneration(LlavaBaseForCausalLM): Returns: torch.Tensor: features from image inputs, concatenated """ + # Requesting hidden states materialises one tensor per layer (~1.8 GiB per + # 1540px image); the plain forward returns the same tensor as the last entry. + last_layer_only = self.vision_feature_layer == -1 features = [] for item in items: # in each item, we assume pixel_values is always batched pixel_values, image_sizes = item.feature, item.image_sizes - image_outputs = self.vision_tower( - pixel_values, image_sizes, output_hidden_states=True - ) - selected_image_feature = image_outputs.hidden_states[ - self.vision_feature_layer - ] + if last_layer_only: + selected_image_feature = self.vision_tower(pixel_values, image_sizes) + else: + image_outputs = self.vision_tower( + pixel_values, image_sizes, output_hidden_states=True + ) + selected_image_feature = image_outputs.hidden_states[ + self.vision_feature_layer + ] if self.vision_feature_select_strategy in ["default", "patch"]: selected_image_feature = selected_image_feature[:, 1:] diff --git a/python/sglang/srt/models/mistral.py b/python/sglang/srt/models/mistral.py index f97ba66bc..dfd9c17e8 100644 --- a/python/sglang/srt/models/mistral.py +++ b/python/sglang/srt/models/mistral.py @@ -115,16 +115,22 @@ class Mistral3ForConditionalGeneration: Returns: torch.Tensor: features from image inputs, concatenated """ + # Requesting hidden states materialises one tensor per layer (~1.8 GiB per + # 1540px image); the plain forward returns the same tensor as the last entry. + last_layer_only = self.vision_feature_layer == -1 features = [] for item in items: # in each item, we assume pixel_values is always batched pixel_values, image_sizes = item.feature, item.image_sizes - image_outputs = self.vision_tower( - pixel_values, image_sizes, output_hidden_states=True - ) - selected_image_feature = image_outputs.hidden_states[ - self.vision_feature_layer - ] + if last_layer_only: + selected_image_feature = self.vision_tower(pixel_values, image_sizes) + else: + image_outputs = self.vision_tower( + pixel_values, image_sizes, output_hidden_states=True + ) + selected_image_feature = image_outputs.hidden_states[ + self.vision_feature_layer + ] if self.vision_feature_select_strategy in ["default", "patch"]: selected_image_feature = selected_image_feature[:, 1:] diff --git a/test/registered/unit/models/test_mistral3_vision_feature.py b/test/registered/unit/models/test_mistral3_vision_feature.py new file mode 100644 index 000000000..c7df5d06b --- /dev/null +++ b/test/registered/unit/models/test_mistral3_vision_feature.py @@ -0,0 +1,109 @@ +"""Pixtral adaptors should not ask the vision tower for every layer to read one. + +The tower materialises one hidden-state tensor per layer when hidden states are +requested, ~49x the tensor the model actually consumes. These tests pin that the +final-layer case takes the cheap path and that both paths agree, for each +adaptor that reads the Pixtral tower this way. +""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.models.llava import LlavaForConditionalGeneration +from sglang.srt.models.mistral import Mistral3ForConditionalGeneration +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + +NUM_LAYERS = 4 +TOKENS = 3 +HIDDEN = 8 + +ADAPTORS = { + "mistral3": Mistral3ForConditionalGeneration.get_image_feature, + "llava": LlavaForConditionalGeneration.get_image_feature, +} + + +class RecordingTower: + """Stands in for PixtralVisionModel, mirroring how it answers. + + ``all_hidden_states`` is the input embedding followed by one tensor per + layer, and the value returned when hidden states are *not* requested is the + last entry of that list -- the property the fix relies on. + """ + + def __init__(self): + torch.manual_seed(0) + self.layers = [torch.randn(1, TOKENS, HIDDEN) for _ in range(NUM_LAYERS + 1)] + self.calls = [] + + def __call__(self, pixel_values, image_sizes, output_hidden_states=False): + self.calls.append(output_hidden_states) + if output_hidden_states: + return SimpleNamespace( + last_hidden_state=self.layers[-1], hidden_states=list(self.layers) + ) + return self.layers[-1] + + +def _model(vision_feature_layer): + return SimpleNamespace( + vision_tower=RecordingTower(), + vision_feature_layer=vision_feature_layer, + vision_feature_select_strategy="full", + multi_modal_projector=lambda feature, image_sizes=None: feature, + ) + + +def _items(n): + return [ + SimpleNamespace(feature=torch.zeros(1, 3, 4, 4), image_sizes=[(4, 4)]) + for _ in range(n) + ] + + +class TestPixtralAdaptorVisionFeature(unittest.TestCase): + def test_final_layer_never_requests_hidden_states(self): + for name, get_image_feature in ADAPTORS.items(): + with self.subTest(adaptor=name): + model = _model(-1) + get_image_feature(model, _items(3)) + self.assertEqual(model.vision_tower.calls, [False, False, False]) + + def test_non_final_layer_still_requests_hidden_states(self): + for name, get_image_feature in ADAPTORS.items(): + with self.subTest(adaptor=name): + model = _model(-2) + get_image_feature(model, _items(2)) + self.assertEqual(model.vision_tower.calls, [True, True]) + + def test_both_paths_agree_on_the_final_layer(self): + # -1 through the cheap path vs NUM_LAYERS (the same tensor, reached by + # indexing the hidden-state list) must produce identical features. + for name, get_image_feature in ADAPTORS.items(): + with self.subTest(adaptor=name): + cheap = get_image_feature(_model(-1), _items(2)) + listed = get_image_feature(_model(NUM_LAYERS), _items(2)) + self.assertTrue(torch.equal(cheap, listed)) + + def test_non_final_layer_selects_that_layer(self): + for name, get_image_feature in ADAPTORS.items(): + with self.subTest(adaptor=name): + model = _model(1) + out = get_image_feature(model, _items(1)) + self.assertTrue( + torch.equal(out, model.vision_tower.layers[1].squeeze(0)) + ) + + def test_features_are_concatenated_per_item(self): + for name, get_image_feature in ADAPTORS.items(): + with self.subTest(adaptor=name): + out = get_image_feature(_model(-1), _items(3)) + self.assertEqual(out.shape, (3 * TOKENS, HIDDEN)) + + +if __name__ == "__main__": + unittest.main()