diff --git a/src/worker/executors/transformers_executor.py b/src/worker/executors/transformers_executor.py index 658c9a7..a18e7d9 100644 --- a/src/worker/executors/transformers_executor.py +++ b/src/worker/executors/transformers_executor.py @@ -114,6 +114,28 @@ logger = logging.getLogger(__name__) +def _select_vision_features( + hidden_states: "Sequence[torch.Tensor]", + feature_layer: int | list[int], + select_strategy: str, +) -> "torch.Tensor": + """Select vision-tower hidden states the way LlavaModel.get_image_features does. + + Crops the leading CLS token when ``select_strategy == "default"``. A + list-valued ``feature_layer`` gathers each layer and concatenates along the + feature dimension. + """ + if isinstance(feature_layer, int): + selected = hidden_states[feature_layer] + if select_strategy == "default": + selected = selected[:, 1:] + return selected + pool = [hidden_states[layer] for layer in feature_layer] + if select_strategy == "default": + pool = [layer[:, 1:] for layer in pool] + return torch.cat(pool, dim=-1) + + class TransformersResult(BaseExecutorResult): ok: bool = True model: str | None = None @@ -458,9 +480,12 @@ def _run_inner( vision_outputs = base_model.vision_tower( inputs.pixel_values, output_hidden_states=True ) - selected_features = vision_outputs.hidden_states[ - self._model.config.vision_feature_layer - ] + model_config = self._model.config + selected_features = _select_vision_features( + vision_outputs.hidden_states, + model_config.vision_feature_layer, + model_config.vision_feature_select_strategy, + ) visual_embeddings = base_model.multi_modal_projector(selected_features) grouped_visual_embeddings: list[torch.Tensor] diff --git a/tests/worker/test_transformers_visual_embedding_features.py b/tests/worker/test_transformers_visual_embedding_features.py new file mode 100644 index 0000000..4289b72 --- /dev/null +++ b/tests/worker/test_transformers_visual_embedding_features.py @@ -0,0 +1,45 @@ +# ruff: noqa: E402 +import pytest + +torch = pytest.importorskip( + "torch", reason="torch not installed (needs --extra inference)" +) + +from worker.executors.transformers_executor import _select_vision_features + + +def _hidden_states(num_layers: int = 3): + return tuple( + torch.arange(2 * 4 * 3, dtype=torch.float32).reshape(2, 4, 3) + layer * 100 + for layer in range(num_layers) + ) + + +def test_int_layer_default_crops_cls_token() -> None: + hidden_states = _hidden_states() + selected = _select_vision_features(hidden_states, -2, "default") + assert selected.shape == (2, 3, 3) + torch.testing.assert_close(selected, hidden_states[-2][:, 1:]) + + +def test_int_layer_full_keeps_cls_token() -> None: + hidden_states = _hidden_states() + selected = _select_vision_features(hidden_states, -2, "full") + assert selected.shape == (2, 4, 3) + torch.testing.assert_close(selected, hidden_states[-2]) + + +def test_list_layer_default_crops_then_concatenates_on_feature_dim() -> None: + hidden_states = _hidden_states() + selected = _select_vision_features(hidden_states, [-2, -1], "default") + assert selected.shape == (2, 3, 6) + expected = torch.cat([hidden_states[-2][:, 1:], hidden_states[-1][:, 1:]], dim=-1) + torch.testing.assert_close(selected, expected) + + +def test_list_layer_full_concatenates_without_cropping() -> None: + hidden_states = _hidden_states() + selected = _select_vision_features(hidden_states, [0, 2], "full") + assert selected.shape == (2, 4, 6) + expected = torch.cat([hidden_states[0], hidden_states[2]], dim=-1) + torch.testing.assert_close(selected, expected)