diff --git a/packages/backend/embedding_atlas/projection.py b/packages/backend/embedding_atlas/projection.py index 1b174708..33bd0353 100644 --- a/packages/backend/embedding_atlas/projection.py +++ b/packages/backend/embedding_atlas/projection.py @@ -321,6 +321,11 @@ def _detect_binary_modality(data: bytes) -> str: def _infer_modality(series: nw.Series) -> str: """Infer the modality by inspecting the first non-null value in the series.""" + # Typed list/array columns (e.g. polars) don't return a Python list per value + dtype = series.dtype + if isinstance(dtype, (nw.List, nw.Array)) and dtype.inner.is_numeric(): + return "vector" + non_null = series.drop_nulls() if len(non_null) == 0: return "text" @@ -332,7 +337,7 @@ def _infer_modality(series: nw.Series) -> str: if ( isinstance(sample, list) and len(sample) > 0 - and isinstance(sample[0], (int, float)) + and isinstance(sample[0], (int, float, np.number)) ): return "vector" diff --git a/packages/backend/tests/test_modality_detection.py b/packages/backend/tests/test_modality_detection.py index 5401f804..ff0c9e9a 100644 --- a/packages/backend/tests/test_modality_detection.py +++ b/packages/backend/tests/test_modality_detection.py @@ -5,6 +5,7 @@ import narwhals as nw import numpy as np import pandas as pd +import polars as pl import pytest from embedding_atlas.projection import ( _detect_binary_modality, @@ -157,6 +158,28 @@ def test_vector_empty_list_falls_through_to_text(self): series = _nw_series([[], []]) assert _infer_modality(series) == "text" + def test_vector_list_of_numpy_scalars(self): + vectors = np.array([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]], dtype=np.float32) + series = _nw_series([list(v) for v in vectors]) + assert _infer_modality(series) == "vector" + + def test_vector_polars_list(self): + series = nw.from_native( + pl.Series([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]), series_only=True + ) + assert _infer_modality(series) == "vector" + + def test_vector_polars_array(self): + series = nw.from_native( + pl.Series([[1.0, 2.0], [3.0, 4.0]], dtype=pl.Array(pl.Float32, 2)), + series_only=True, + ) + assert _infer_modality(series) == "vector" + + def test_polars_list_of_strings_is_text(self): + series = nw.from_native(pl.Series([["a", "b"], ["c"]]), series_only=True) + assert _infer_modality(series) == "text" + def test_image_bytes(self): png_bytes = b"\x89PNG\r\n\x1a\n" + b"\x00" * 20 series = _nw_series([png_bytes, png_bytes])