From 00e4524bb47cfd43098d0a0128a867ead10374fc Mon Sep 17 00:00:00 2001 From: LeonardoRosaa Date: Mon, 13 Jul 2026 11:37:33 -0300 Subject: [PATCH 1/4] add `image_resolver.get_for_export` with collection filter support --- .../resolvers/image_resolver/__init__.py | 2 + .../get_all_by_collection_id.py | 15 +- .../image_resolver/get_for_export.py | 56 +++++++ .../image_resolver/test_get_for_export.py | 158 ++++++++++++++++++ 4 files changed, 227 insertions(+), 4 deletions(-) create mode 100644 lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py create mode 100644 lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py diff --git a/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py b/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py index b219a39846..12f81ab07f 100644 --- a/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py +++ b/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py @@ -11,6 +11,7 @@ ) from lightly_studio.resolvers.image_resolver.get_by_id import get_by_id from lightly_studio.resolvers.image_resolver.get_dimension_bounds import get_dimension_bounds +from lightly_studio.resolvers.image_resolver.get_for_export import get_for_export from lightly_studio.resolvers.image_resolver.get_many_by_id import get_many_by_id from lightly_studio.resolvers.image_resolver.get_sample_ids import ( build_sample_ids_query, @@ -30,6 +31,7 @@ "get_all_by_collection_id", "get_by_id", "get_dimension_bounds", + "get_for_export", "get_many_by_id", "get_sample_ids", "get_sample_ids_by_paths", diff --git a/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_all_by_collection_id.py b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_all_by_collection_id.py index e5988a6726..d5a52eb744 100644 --- a/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_all_by_collection_id.py +++ b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_all_by_collection_id.py @@ -100,13 +100,20 @@ def get_all_by_collection_id( # noqa: PLR0913 """Retrieve samples for a specific collection with optional filtering.""" # Resolve any embedding-plot region selection to concrete sample ids on the filter before the # query is built (the point-in-polygon test needs the session, which `apply` lacks). - sample_filter = filters.sample_filter if filters is not None else None - if sample_filter is not None and sample_filter.embedding_region is not None: - sample_filter.region_sample_ids = embedding_region_resolver.get_sample_ids_in_region( + if ( + filters is not None + and filters.sample_filter is not None + and filters.sample_filter.embedding_region is not None + ): + region_sample_ids = embedding_region_resolver.get_sample_ids_in_region( session=session, collection_id=collection_id, - region=sample_filter.embedding_region, + region=filters.sample_filter.embedding_region, ) + resolved_sample_filter = filters.sample_filter.model_copy( + update={"region_sample_ids": region_sample_ids} + ) + filters = filters.model_copy(update={"sample_filter": resolved_sample_filter}) embedding_model_id, distance_expr = get_distance_expression( session=session, diff --git a/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py new file mode 100644 index 0000000000..dfb9802338 --- /dev/null +++ b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py @@ -0,0 +1,56 @@ +"""Implementation of get_for_export function for images.""" + +from __future__ import annotations + +from collections.abc import Generator +from uuid import UUID + +from sqlmodel import Session, col, select + +from lightly_studio.core.image.image_sample import ImageSample +from lightly_studio.models.image import ImageTable +from lightly_studio.models.sample import SampleTable +from lightly_studio.resolvers import embedding_region_resolver +from lightly_studio.resolvers.image_filter import ImageFilter + + +def get_for_export( + session: Session, + collection_id: UUID, + collection_filter: ImageFilter | None, +) -> Generator[ImageSample, None, None]: + """Return all images in a collection as a lazy generator of ImageSamples. + + If ``collection_filter`` is provided, only images matching the filter are + returned. Embedding-region filters are resolved to concrete sample IDs + before the query is executed. + + Args: + session: Database session. + collection_id: ID of the collection to export. + collection_filter: Optional filter to restrict which images are returned. + + Returns: + Generator of ImageSamples for the matching images. + """ + query = ( + select(ImageTable).join(ImageTable.sample).where(SampleTable.collection_id == collection_id) + ) + if collection_filter is not None: + sample_filter = collection_filter.sample_filter + if sample_filter is not None and sample_filter.embedding_region is not None: + region_sample_ids = embedding_region_resolver.get_sample_ids_in_region( + session=session, + collection_id=collection_id, + region=sample_filter.embedding_region, + ) + resolved_sample_filter = sample_filter.model_copy( + update={"region_sample_ids": region_sample_ids} + ) + collection_filter = collection_filter.model_copy( + update={"sample_filter": resolved_sample_filter} + ) + query = query.where( + col(SampleTable.sample_id).in_(collection_filter.build_sample_ids_query(collection_id)) + ) + return (ImageSample(row) for row in session.exec(query)) diff --git a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py new file mode 100644 index 0000000000..5f2e679606 --- /dev/null +++ b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py @@ -0,0 +1,158 @@ +from __future__ import annotations + +from sqlmodel import Session + +from lightly_studio.models.embedding_region import EmbeddingRegion, Point2D +from lightly_studio.models.two_dim_embedding import TwoDimEmbeddingTable +from lightly_studio.resolvers import image_resolver, sample_embedding_resolver +from lightly_studio.resolvers.image_filter import FilterDimensions, ImageFilter +from lightly_studio.resolvers.sample_resolver.sample_filter import SampleFilter +from tests.helpers_resolvers import ( + create_collection, + create_embedding_model, + create_image, + create_sample_embedding, +) + + +def test_get_for_export__no_filter(db_session: Session) -> None: + collection = create_collection(session=db_session) + other_collection = create_collection(session=db_session) + + image1 = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/img1.jpg", + ) + image2 = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/img2.jpg", + ) + create_image( + session=db_session, + collection_id=other_collection.collection_id, + file_path_abs="/data/other.jpg", + ) + + result = list( + image_resolver.get_for_export( + session=db_session, + collection_id=collection.collection_id, + collection_filter=None, + ) + ) + + assert {s.sample_id for s in result} == {image1.sample_id, image2.sample_id} + + +def test_get_for_export__with_image_filter(db_session: Session) -> None: + collection = create_collection(session=db_session) + + small_image = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/small.jpg", + width=100, + height=100, + ) + large_image = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/large.jpg", + width=800, + height=800, + ) + + result = list( + image_resolver.get_for_export( + session=db_session, + collection_id=collection.collection_id, + collection_filter=ImageFilter(width=FilterDimensions(min=500)), + ) + ) + + assert {s.sample_id for s in result} == {large_image.sample_id} + + +def test_get_for_export__with_embedding_region_filter(db_session: Session) -> None: + collection = create_collection(session=db_session) + embedding_model = create_embedding_model( + session=db_session, + collection_id=collection.collection_id, + embedding_dimension=3, + ) + + image_inside1 = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="inside1.jpg", + ) + create_sample_embedding( + session=db_session, + sample_id=image_inside1.sample_id, + embedding_model_id=embedding_model.embedding_model_id, + embedding=[1.0, 0.2, 0.3], + ) + image_inside2 = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="inside2.jpg", + ) + create_sample_embedding( + session=db_session, + sample_id=image_inside2.sample_id, + embedding_model_id=embedding_model.embedding_model_id, + embedding=[2.0, 0.2, 0.3], + ) + image_outside = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="outside.jpg", + ) + create_sample_embedding( + session=db_session, + sample_id=image_outside.sample_id, + embedding_model_id=embedding_model.embedding_model_id, + embedding=[3.0, 0.2, 0.3], + ) + + # Seed 2D coordinates: inside1 and inside2 within (0,0)-(10,10), outside at (100, 100). + cache_key, sample_ids_in_order = sample_embedding_resolver.get_hash_by_collection_id( + session=db_session, + collection_id=collection.collection_id, + embedding_model_id=embedding_model.embedding_model_id, + ) + inside1_i = sample_ids_in_order.index(image_inside1.sample_id) + inside2_i = sample_ids_in_order.index(image_inside2.sample_id) + outside_i = sample_ids_in_order.index(image_outside.sample_id) + x = [0.0, 0.0, 0.0] + y = [0.0, 0.0, 0.0] + x[inside1_i] = 1.0 + y[inside1_i] = 1.0 + x[inside2_i] = 5.0 + y[inside2_i] = 5.0 + x[outside_i] = 100.0 + y[outside_i] = 100.0 + db_session.add(TwoDimEmbeddingTable(hash=cache_key, x=x, y=y)) + db_session.commit() + + region = EmbeddingRegion( + polygon=[ + Point2D(x=0, y=0), + Point2D(x=10, y=0), + Point2D(x=10, y=10), + Point2D(x=0, y=10), + ] + ) + collection_filter = ImageFilter(sample_filter=SampleFilter(embedding_region=region)) + + result = list( + image_resolver.get_for_export( + session=db_session, + collection_id=collection.collection_id, + collection_filter=collection_filter, + ) + ) + + assert {s.sample_id for s in result} == {image_inside1.sample_id, image_inside2.sample_id} From 30cc00a8355c6b4ba7747c49dc1318f1d1b0e600 Mon Sep 17 00:00:00 2001 From: LeonardoRosaa Date: Mon, 13 Jul 2026 11:50:03 -0300 Subject: [PATCH 2/4] format file --- .../tests/resolvers/image_resolver/test_get_for_export.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py index 5f2e679606..310cd8ca8f 100644 --- a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py +++ b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py @@ -49,7 +49,7 @@ def test_get_for_export__no_filter(db_session: Session) -> None: def test_get_for_export__with_image_filter(db_session: Session) -> None: collection = create_collection(session=db_session) - small_image = create_image( + create_image( session=db_session, collection_id=collection.collection_id, file_path_abs="/data/small.jpg", From d488a54aef3c6dd21b7c903eeeb1ddbcdf20d1f9 Mon Sep 17 00:00:00 2001 From: LeonardoRosaa Date: Wed, 15 Jul 2026 07:45:39 -0300 Subject: [PATCH 3/4] fix n+1 --- .../resolvers/image_resolver/__init__.py | 6 ++- .../image_resolver/get_for_export.py | 40 +++++++++++++++- .../image_resolver/test_get_for_export.py | 46 +++++++++++++++++++ 3 files changed, 90 insertions(+), 2 deletions(-) diff --git a/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py b/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py index 12f81ab07f..2238f77c08 100644 --- a/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py +++ b/lightly_studio/src/lightly_studio/resolvers/image_resolver/__init__.py @@ -11,7 +11,10 @@ ) from lightly_studio.resolvers.image_resolver.get_by_id import get_by_id from lightly_studio.resolvers.image_resolver.get_dimension_bounds import get_dimension_bounds -from lightly_studio.resolvers.image_resolver.get_for_export import get_for_export +from lightly_studio.resolvers.image_resolver.get_for_export import ( + ImageExportPreload, + get_for_export, +) from lightly_studio.resolvers.image_resolver.get_many_by_id import get_many_by_id from lightly_studio.resolvers.image_resolver.get_sample_ids import ( build_sample_ids_query, @@ -23,6 +26,7 @@ from lightly_studio.resolvers.image_resolver.get_samples_excluding import get_samples_excluding __all__ = [ + "ImageExportPreload", "build_sample_ids_query", "count_image_annotations_by_collection", "create_many", diff --git a/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py index dfb9802338..209ae429e6 100644 --- a/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py +++ b/lightly_studio/src/lightly_studio/resolvers/image_resolver/get_for_export.py @@ -3,21 +3,37 @@ from __future__ import annotations from collections.abc import Generator +from enum import Enum from uuid import UUID +from sqlalchemy.orm import contains_eager, joinedload, selectinload +from sqlalchemy.orm.interfaces import LoaderOption from sqlmodel import Session, col, select from lightly_studio.core.image.image_sample import ImageSample +from lightly_studio.models.annotation.annotation_base import AnnotationBaseTable from lightly_studio.models.image import ImageTable from lightly_studio.models.sample import SampleTable from lightly_studio.resolvers import embedding_region_resolver from lightly_studio.resolvers.image_filter import ImageFilter +class ImageExportPreload(Enum): + """Relationships to eagerly load when exporting images. + + Pass a frozenset of these values to ``get_for_export`` to avoid N+1 queries + when accessing the corresponding properties on the returned ``ImageSample``s. + """ + + ANNOTATIONS = "annotations" + CAPTIONS = "captions" + + def get_for_export( session: Session, collection_id: UUID, collection_filter: ImageFilter | None, + preload: frozenset[ImageExportPreload] = frozenset(), ) -> Generator[ImageSample, None, None]: """Return all images in a collection as a lazy generator of ImageSamples. @@ -29,12 +45,34 @@ def get_for_export( session: Database session. collection_id: ID of the collection to export. collection_filter: Optional filter to restrict which images are returned. + preload: Set of relationships to eagerly load. By default nothing is + preloaded. Pass ``frozenset({ImageExportPreload.ANNOTATIONS})`` or + ``frozenset({ImageExportPreload.CAPTIONS})``. Returns: Generator of ImageSamples for the matching images. """ + sample_options: list[LoaderOption] = [] + if ImageExportPreload.ANNOTATIONS in preload: + sample_options.append( + selectinload(SampleTable.annotations).options( + joinedload(AnnotationBaseTable.annotation_label), + joinedload(AnnotationBaseTable.object_detection_details), + joinedload(AnnotationBaseTable.segmentation_details), + ) + ) + if ImageExportPreload.CAPTIONS in preload: + sample_options.append(selectinload(SampleTable.captions)) + + eager_sample = contains_eager(ImageTable.sample) + if sample_options: + eager_sample = eager_sample.options(*sample_options) # type: ignore[arg-type] + query = ( - select(ImageTable).join(ImageTable.sample).where(SampleTable.collection_id == collection_id) + select(ImageTable) + .join(ImageTable.sample) + .options(eager_sample) + .where(SampleTable.collection_id == collection_id) ) if collection_filter is not None: sample_filter = collection_filter.sample_filter diff --git a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py index 310cd8ca8f..f6dc0bc680 100644 --- a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py +++ b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py @@ -6,8 +6,12 @@ from lightly_studio.models.two_dim_embedding import TwoDimEmbeddingTable from lightly_studio.resolvers import image_resolver, sample_embedding_resolver from lightly_studio.resolvers.image_filter import FilterDimensions, ImageFilter +from lightly_studio.resolvers.image_resolver import ImageExportPreload from lightly_studio.resolvers.sample_resolver.sample_filter import SampleFilter from tests.helpers_resolvers import ( + create_annotation, + create_annotation_label, + create_caption, create_collection, create_embedding_model, create_image, @@ -156,3 +160,45 @@ def test_get_for_export__with_embedding_region_filter(db_session: Session) -> No ) assert {s.sample_id for s in result} == {image_inside1.sample_id, image_inside2.sample_id} + + +def test_get_for_export__preloaded_data_accessible(db_session: Session) -> None: + collection = create_collection(session=db_session) + image = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/img.jpg", + ) + label = create_annotation_label( + session=db_session, + root_collection_id=collection.collection_id, + label_name="cat", + ) + create_annotation( + session=db_session, + collection_id=collection.collection_id, + sample_id=image.sample_id, + annotation_label_id=label.annotation_label_id, + annotation_data={"x": 10, "y": 10, "width": 20, "height": 20}, + ) + create_caption( + session=db_session, + collection_id=collection.collection_id, + parent_sample_id=image.sample_id, + text="a cat sitting on a mat", + ) + + result = list( + image_resolver.get_for_export( + session=db_session, + collection_id=collection.collection_id, + collection_filter=None, + preload=frozenset({ImageExportPreload.ANNOTATIONS, ImageExportPreload.CAPTIONS}), + ) + ) + + assert len(result) == 1 + annotations = result[0].annotations + assert len(annotations) == 1 + assert annotations[0].class_name == "cat" + assert result[0].captions == ["a cat sitting on a mat"] From 70e368fc6752bbfc6b8918ad3dead991306cc153 Mon Sep 17 00:00:00 2001 From: LeonardoRosaa Date: Wed, 15 Jul 2026 10:32:46 -0300 Subject: [PATCH 4/4] validate N+1 --- .../image_resolver/test_get_for_export.py | 52 +++++++++++++++++++ 1 file changed, 52 insertions(+) diff --git a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py index f6dc0bc680..3ef9e4d17a 100644 --- a/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py +++ b/lightly_studio/tests/resolvers/image_resolver/test_get_for_export.py @@ -1,5 +1,6 @@ from __future__ import annotations +from sqlalchemy import inspect from sqlmodel import Session from lightly_studio.models.embedding_region import EmbeddingRegion, Point2D @@ -162,6 +163,48 @@ def test_get_for_export__with_embedding_region_filter(db_session: Session) -> No assert {s.sample_id for s in result} == {image_inside1.sample_id, image_inside2.sample_id} +def test_get_for_export__without_preload_relationships_are_unloaded(db_session: Session) -> None: + collection = create_collection(session=db_session) + image = create_image( + session=db_session, + collection_id=collection.collection_id, + file_path_abs="/data/img.jpg", + ) + label = create_annotation_label( + session=db_session, + root_collection_id=collection.collection_id, + label_name="cat", + ) + create_annotation( + session=db_session, + collection_id=collection.collection_id, + sample_id=image.sample_id, + annotation_label_id=label.annotation_label_id, + annotation_data={"x": 10, "y": 10, "width": 20, "height": 20}, + ) + create_caption( + session=db_session, + collection_id=collection.collection_id, + parent_sample_id=image.sample_id, + text="a cat sitting on a mat", + ) + db_session.expire_all() + + result = list( + image_resolver.get_for_export( + session=db_session, + collection_id=collection.collection_id, + collection_filter=None, + ) + ) + + assert len(result) == 1 + sample_state = inspect(result[0].sample_table) + assert sample_state is not None + assert "annotations" in sample_state.unloaded + assert "captions" in sample_state.unloaded + + def test_get_for_export__preloaded_data_accessible(db_session: Session) -> None: collection = create_collection(session=db_session) image = create_image( @@ -187,6 +230,7 @@ def test_get_for_export__preloaded_data_accessible(db_session: Session) -> None: parent_sample_id=image.sample_id, text="a cat sitting on a mat", ) + db_session.expire_all() result = list( image_resolver.get_for_export( @@ -198,6 +242,14 @@ def test_get_for_export__preloaded_data_accessible(db_session: Session) -> None: ) assert len(result) == 1 + sample_state = inspect(result[0].sample_table) + assert sample_state is not None + assert "annotations" not in sample_state.unloaded + assert "captions" not in sample_state.unloaded + annotation_state = inspect(result[0].sample_table.annotations[0]) + assert annotation_state is not None + assert "annotation_label" not in annotation_state.unloaded + assert "object_detection_details" not in annotation_state.unloaded annotations = result[0].annotations assert len(annotations) == 1 assert annotations[0].class_name == "cat"