Fix Mistral3 retaining every vision-tower layer to read one (#39185)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
Alison Shao
2026-09-18 16:11:54 -07:00
committed by GitHub
co-authored by Xinyuan Tong Xinyuan Tong
parent fa521e2758
commit 2394b231c2
3 changed files with 133 additions and 12 deletions
+12 -6
View File
@@ -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:]
+12 -6
View File
@@ -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:]