Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 28 additions & 3 deletions src/worker/executors/transformers_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand Down
45 changes: 45 additions & 0 deletions tests/worker/test_transformers_visual_embedding_features.py
Original file line number Diff line number Diff line change
@@ -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)