From a992f96f220099c66ca367ba6ce822869ad02d71 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20K=C3=A4nzig?= <36882833+nkaenzig@users.noreply.github.com> Date: Thu, 28 Aug 2025 12:31:10 +0200 Subject: [PATCH 1/3] Add `eva.multimodal` (#869) --- README.md | 5 +- .../online/multiple_choice}/pubmedqa.yaml | 16 +- .../multiple_choice/patch_camelyon.yaml | 63 ++++++ .../multiple_choice/patch_camelyon.yaml | 65 +++++++ docs/index.md | 7 +- docs/reference/language/models/networks.md | 2 +- docs/reference/language/models/wrappers.md | 6 +- docs/reference/multimodal/data/datasets.md | 6 + docs/reference/multimodal/data/index.md | 3 + docs/reference/multimodal/index.md | 8 + docs/reference/multimodal/models/modules.md | 5 + docs/reference/multimodal/models/networks.md | 14 ++ docs/reference/multimodal/models/wrappers.md | 8 + docs/reference/multimodal/utils/image.md | 5 + .../tutorials/pubmedqa_classification.md | 23 +-- mkdocs.yml | 11 ++ pdm.lock | 13 +- pyproject.toml | 6 + src/eva/core/cli/setup.py | 2 +- src/eva/core/data/dataloaders/__init__.py | 3 +- .../data/dataloaders/collate_fn/__init__.py | 5 - .../data/dataloaders/collate_fn/collate.py | 24 --- src/eva/core/models/wrappers/base.py | 4 +- src/eva/core/models/wrappers/from_function.py | 6 +- src/eva/core/models/wrappers/from_torchhub.py | 16 +- src/eva/core/models/wrappers/huggingface.py | 9 +- src/eva/core/models/wrappers/onnx.py | 10 +- src/eva/language/__init__.py | 3 +- src/eva/language/data/dataloaders/__init__.py | 5 + .../data/dataloaders/collate_fn/__init__.py | 5 + .../data/dataloaders/collate_fn/text.py | 32 ++++ src/eva/language/data/datasets/__init__.py | 2 +- .../data/datasets/{language.py => base.py} | 2 +- .../data/datasets/classification/base.py | 46 +---- .../data/datasets/classification/pubmedqa.py | 12 +- src/eva/language/data/datasets/schemas.py | 15 ++ src/eva/language/data/datasets/text.py | 93 +++++++++ src/eva/language/data/datasets/typings.py | 23 +++ src/eva/language/data/messages.py | 51 +++++ src/eva/language/models/__init__.py | 24 +-- src/eva/language/models/modules/__init__.py | 4 +- src/eva/language/models/modules/language.py | 55 ++++++ src/eva/language/models/modules/text.py | 85 --------- src/eva/language/models/networks/__init__.py | 12 ++ src/eva/language/models/networks/alibaba.py | 26 +++ .../language/models/networks/api/__init__.py | 11 ++ .../language/models/networks/api/anthropic.py | 34 ++++ src/eva/language/models/networks/registry.py | 5 + src/eva/language/models/typings.py | 23 +++ src/eva/language/models/wrappers/__init__.py | 11 +- src/eva/language/models/wrappers/base.py | 47 +++++ .../language/models/wrappers/from_registry.py | 54 ++++++ .../language/models/wrappers/huggingface.py | 52 ++++- src/eva/language/models/wrappers/litellm.py | 127 +++++++----- src/eva/language/models/wrappers/vllm.py | 50 +++-- src/eva/language/utils/__init__.py | 3 +- src/eva/language/utils/str_to_int_tensor.py | 27 +-- src/eva/language/utils/text/__init__.py | 5 + src/eva/language/utils/text/messages.py | 67 +++++++ src/eva/multimodal/__init__.py | 6 + src/eva/multimodal/data/__init__.py | 5 + .../multimodal/data/dataloaders/__init__.py | 5 + .../data/dataloaders/collate_fn/__init__.py | 5 + .../data/dataloaders/collate_fn/text_image.py | 28 +++ src/eva/multimodal/data/datasets/__init__.py | 6 + src/eva/multimodal/data/datasets/base.py | 13 ++ .../data/datasets/multiple_choice/__init__.py | 5 + .../multiple_choice/patch_camelyon.py | 80 ++++++++ src/eva/multimodal/data/datasets/schemas.py | 14 ++ .../multimodal/data/datasets/text_image.py | 77 ++++++++ src/eva/multimodal/data/datasets/typings.py | 27 +++ src/eva/multimodal/models/__init__.py | 8 + src/eva/multimodal/models/modules/__init__.py | 5 + .../models/modules/vision_language.py | 55 ++++++ .../multimodal/models/networks/__init__.py | 14 ++ src/eva/multimodal/models/networks/alibaba.py | 39 ++++ .../models/networks/api/__init__.py | 11 ++ .../models/networks/api/anthropic.py | 34 ++++ src/eva/multimodal/models/networks/others.py | 47 +++++ .../multimodal/models/networks/registry.py | 5 + src/eva/multimodal/models/typings.py | 27 +++ .../multimodal/models/wrappers/__init__.py | 13 ++ src/eva/multimodal/models/wrappers/base.py | 47 +++++ .../models/wrappers/from_registry.py | 54 ++++++ .../multimodal/models/wrappers/huggingface.py | 180 ++++++++++++++++++ src/eva/multimodal/models/wrappers/litellm.py | 56 ++++++ src/eva/multimodal/utils/__init__.py | 1 + src/eva/multimodal/utils/image/__init__.py | 5 + src/eva/multimodal/utils/image/encode.py | 28 +++ src/eva/multimodal/utils/text/__init__.py | 1 + src/eva/multimodal/utils/text/messages.py | 79 ++++++++ .../datasets/classification/patch_camelyon.py | 14 +- src/eva/vision/data/transforms/__init__.py | 3 +- .../data/transforms/spatial/__init__.py | 3 +- .../transforms/spatial/functional/__init__.py | 5 + .../transforms/spatial/functional/resize.py | 26 +++ .../vision/data/transforms/spatial/resize.py | 62 ++++++ .../vision/models/wrappers/from_registry.py | 11 +- src/eva/vision/models/wrappers/from_timm.py | 10 +- tests/__init__.py | 2 +- tests/eva/core/test_cli.py | 2 +- tests/eva/language/__init__.py | 2 +- .../datasets/classification/test_pubmedqa.py | 11 +- .../language/models/modules/test_language.py | 62 ++++++ .../eva/language/models/modules/test_text.py | 69 ------- .../models/wrappers/test_huggingface.py | 11 +- .../language/models/wrappers/test_litellm.py | 14 +- .../eva/language/models/wrappers/test_vllm.py | 16 +- tests/eva/language/test_language_cli.py | 7 +- .../language/utils/test_str_to_int_tensor.py | 18 +- tests/eva/multimodal/__init__.py | 1 + tests/eva/multimodal/test_multimodal_cli.py | 72 +++++++ tests/eva/vision/__init__.py | 2 +- 113 files changed, 2322 insertions(+), 437 deletions(-) rename configs/language/{ => pathology/online/multiple_choice}/pubmedqa.yaml (72%) create mode 100644 configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml create mode 100644 configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml create mode 100644 docs/reference/multimodal/data/datasets.md create mode 100644 docs/reference/multimodal/data/index.md create mode 100644 docs/reference/multimodal/index.md create mode 100644 docs/reference/multimodal/models/modules.md create mode 100644 docs/reference/multimodal/models/networks.md create mode 100644 docs/reference/multimodal/models/wrappers.md create mode 100644 docs/reference/multimodal/utils/image.md delete mode 100644 src/eva/core/data/dataloaders/collate_fn/__init__.py delete mode 100644 src/eva/core/data/dataloaders/collate_fn/collate.py create mode 100644 src/eva/language/data/dataloaders/__init__.py create mode 100644 src/eva/language/data/dataloaders/collate_fn/__init__.py create mode 100644 src/eva/language/data/dataloaders/collate_fn/text.py rename src/eva/language/data/datasets/{language.py => base.py} (84%) create mode 100644 src/eva/language/data/datasets/schemas.py create mode 100644 src/eva/language/data/datasets/text.py create mode 100644 src/eva/language/data/datasets/typings.py create mode 100644 src/eva/language/data/messages.py create mode 100644 src/eva/language/models/modules/language.py delete mode 100644 src/eva/language/models/modules/text.py create mode 100644 src/eva/language/models/networks/__init__.py create mode 100644 src/eva/language/models/networks/alibaba.py create mode 100644 src/eva/language/models/networks/api/__init__.py create mode 100644 src/eva/language/models/networks/api/anthropic.py create mode 100644 src/eva/language/models/networks/registry.py create mode 100644 src/eva/language/models/typings.py create mode 100644 src/eva/language/models/wrappers/base.py create mode 100644 src/eva/language/models/wrappers/from_registry.py create mode 100644 src/eva/language/utils/text/__init__.py create mode 100644 src/eva/language/utils/text/messages.py create mode 100644 src/eva/multimodal/__init__.py create mode 100644 src/eva/multimodal/data/__init__.py create mode 100644 src/eva/multimodal/data/dataloaders/__init__.py create mode 100644 src/eva/multimodal/data/dataloaders/collate_fn/__init__.py create mode 100644 src/eva/multimodal/data/dataloaders/collate_fn/text_image.py create mode 100644 src/eva/multimodal/data/datasets/__init__.py create mode 100644 src/eva/multimodal/data/datasets/base.py create mode 100644 src/eva/multimodal/data/datasets/multiple_choice/__init__.py create mode 100644 src/eva/multimodal/data/datasets/multiple_choice/patch_camelyon.py create mode 100644 src/eva/multimodal/data/datasets/schemas.py create mode 100644 src/eva/multimodal/data/datasets/text_image.py create mode 100644 src/eva/multimodal/data/datasets/typings.py create mode 100644 src/eva/multimodal/models/__init__.py create mode 100644 src/eva/multimodal/models/modules/__init__.py create mode 100644 src/eva/multimodal/models/modules/vision_language.py create mode 100644 src/eva/multimodal/models/networks/__init__.py create mode 100644 src/eva/multimodal/models/networks/alibaba.py create mode 100644 src/eva/multimodal/models/networks/api/__init__.py create mode 100644 src/eva/multimodal/models/networks/api/anthropic.py create mode 100644 src/eva/multimodal/models/networks/others.py create mode 100644 src/eva/multimodal/models/networks/registry.py create mode 100644 src/eva/multimodal/models/typings.py create mode 100644 src/eva/multimodal/models/wrappers/__init__.py create mode 100644 src/eva/multimodal/models/wrappers/base.py create mode 100644 src/eva/multimodal/models/wrappers/from_registry.py create mode 100644 src/eva/multimodal/models/wrappers/huggingface.py create mode 100644 src/eva/multimodal/models/wrappers/litellm.py create mode 100644 src/eva/multimodal/utils/__init__.py create mode 100644 src/eva/multimodal/utils/image/__init__.py create mode 100644 src/eva/multimodal/utils/image/encode.py create mode 100644 src/eva/multimodal/utils/text/__init__.py create mode 100644 src/eva/multimodal/utils/text/messages.py create mode 100644 src/eva/vision/data/transforms/spatial/functional/__init__.py create mode 100644 src/eva/vision/data/transforms/spatial/functional/resize.py create mode 100644 src/eva/vision/data/transforms/spatial/resize.py create mode 100644 tests/eva/language/models/modules/test_language.py delete mode 100644 tests/eva/language/models/modules/test_text.py create mode 100644 tests/eva/multimodal/__init__.py create mode 100644 tests/eva/multimodal/test_multimodal_cli.py diff --git a/README.md b/README.md index 034f34dc5..7a4f7311a 100644 --- a/README.md +++ b/README.md @@ -34,7 +34,7 @@ Check out the [documentation](https://kaiko-ai.github.io/eva/) for more informat ### Highlights: - Easy and reliable benchmark of Oncology FMs -- Supports patch-level classification, slide-level classification, semantic segmentation, and text classification downstream tasks +- Supports patch-level classification, slide-level classification, semantic segmentation, and (visual) question answering tasks. - Automatic embedding inference and evaluation of a downstream task - Native support of popular medical [datasets](https://kaiko-ai.github.io/eva/dev/datasets/) and models - Produce statistics over multiple evaluation fits and multiple metrics @@ -52,6 +52,9 @@ pip install 'kaiko-eva[vision]' # to install the expanded `language` version pip install 'kaiko-eva[language]' +# to install the expanded `multimodal` version +pip install 'kaiko-eva[multimodal]' + # to install everything pip install 'kaiko-eva[all]' ``` diff --git a/configs/language/pubmedqa.yaml b/configs/language/pathology/online/multiple_choice/pubmedqa.yaml similarity index 72% rename from configs/language/pubmedqa.yaml rename to configs/language/pathology/online/multiple_choice/pubmedqa.yaml index b7eac80f6..cd322b82a 100644 --- a/configs/language/pubmedqa.yaml +++ b/configs/language/pathology/online/multiple_choice/pubmedqa.yaml @@ -6,16 +6,13 @@ trainer: default_root_dir: ${oc.env:OUTPUT_ROOT, logs/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/pubmedqa} checkpoint_type: null model: - class_path: eva.language.models.TextModule + class_path: eva.language.models.LanguageModule init_args: - prompt: "Instruction: Carefully read the question and the provided context. Answer with one word: 'yes', 'no', or 'maybe'. Answer: " model: - class_path: eva.language.models.LiteLLMTextModel + class_path: eva.language.models.wrappers.ModelFromRegistry init_args: - model_name_or_path: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-latest} - model_kwargs: - temperature: 0.0 # should be strictly positive for HF models - # max_new_tokens: 1 # used for HF + model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} + model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} metrics: common: - class_path: eva.metrics.MulticlassClassificationMetrics @@ -25,6 +22,9 @@ model: postprocess: predictions_transforms: - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor + init_args: + mapping: {"no": 0, "yes": 1, "maybe": 2} + case_sensitive: false data: class_path: eva.DataModule init_args: @@ -44,4 +44,4 @@ data: batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 16} num_workers: &N_DATA_WORKERS ${oc.env:N_DATA_WORKERS, 1} shuffle: false - collate_fn: eva.core.data.dataloaders.text_collate_fn \ No newline at end of file + collate_fn: eva.language.data.dataloaders.text_collate diff --git a/configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml b/configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml new file mode 100644 index 000000000..6b2e3675d --- /dev/null +++ b/configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml @@ -0,0 +1,63 @@ +trainer: + class_path: eva.Trainer + init_args: + accelerator: ${oc.env:ACCELERATOR, auto} + n_runs: &N_RUNS ${oc.env:N_RUNS, 2} + default_root_dir: ${oc.env:OUTPUT_ROOT, logs/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/pubmedqa} + precision: bf16 + checkpoint_type: null + callbacks: + - class_path: eva.callbacks.ConfigurationLogger +model: + class_path: eva.multimodal.models.modules.VisionLanguageModule + init_args: + model: + class_path: eva.multimodal.models.wrappers.ModelFromRegistry + init_args: + model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} + model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} + metrics: + common: + - class_path: eva.metrics.MulticlassClassificationMetrics + init_args: + num_classes: 2 + input_type: "discrete" + postprocess: + predictions_transforms: + - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor + init_args: + mapping: {"A": 0, "B": 1} + case_sensitive: false +data: + class_path: eva.DataModule + init_args: + datasets: + val: + class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: &DATASET_ARGS + root: ${oc.env:DATA_ROOT, /mnt/localdisk/data/patch_camelyon} + split: val + download: ${oc.env:DOWNLOAD_DATA, false} + # Set `download: true` to download the dataset from https://zenodo.org/records/1494286 + # The PatchCamelyon dataset is distributed under the following license: + # "Creative Commons Zero v1.0 Universal" + # (see: https://choosealicense.com/licenses/cc0-1.0/) + transforms: + image: + class_path: eva.vision.data.transforms.Resize + init_args: + size: ${oc.env:RESIZE_DIM, null} + max_bytes: ${oc.env:IMAGE_MAX_BYTES, null} + max_samples: 500 + test: + class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: + <<: *DATASET_ARGS + split: test + dataloaders: + val: &DATALOADER_ARGS + batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 16} + num_workers: &N_DATA_WORKERS ${oc.env:N_DATA_WORKERS, 1} + collate_fn: eva.multimodal.data.dataloaders.text_image_collate + test: + <<: *DATALOADER_ARGS diff --git a/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml b/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml new file mode 100644 index 000000000..15bdee4ee --- /dev/null +++ b/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml @@ -0,0 +1,65 @@ +trainer: + class_path: eva.Trainer + init_args: + accelerator: cpu + n_runs: &N_RUNS ${oc.env:N_RUNS, 2} + default_root_dir: &LIGHTNING_ROOT ${oc.env:LIGHTNING_ROOT, logs/test/multimodal/online/patch_camelyon} + max_epochs: &MAX_EPOCHS 1 + limit_val_batches: 2 + limit_test_batches: 2 + precision: bf16 + checkpoint_type: null + callbacks: + - class_path: eva.callbacks.ConfigurationLogger +model: + class_path: eva.multimodal.models.modules.VisionLanguageModule + init_args: + model: + class_path: eva.multimodal.models.wrappers.ModelFromRegistry + init_args: + model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} + model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} + metrics: + common: + - class_path: eva.metrics.MulticlassClassificationMetrics + init_args: + num_classes: 2 + input_type: "discrete" + postprocess: + predictions_transforms: + - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor + init_args: + mapping: {"A": 0, "B": 1} + case_sensitive: false +data: + class_path: eva.DataModule + init_args: + datasets: + val: + class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: &DATASET_ARGS + root: ${oc.env:TESTS_ROOT, tests/eva}/assets/vision/datasets/patch_camelyon + split: val + download: false + transforms: + image: + class_path: eva.vision.data.transforms.Resize + init_args: + size: ${oc.env:RESIZE_DIM, null} + max_bytes: ${oc.env:IMAGE_MAX_BYTES, null} + max_samples: null + test: + class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: + <<: *DATASET_ARGS + split: test + dataloaders: + val: &DATALOADER_ARGS + batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 16} + collate_fn: eva.multimodal.data.dataloaders.text_image_collate + num_workers: 0 + pin_memory: false + persistent_workers: false + prefetch_factor: null + test: + <<: *DATALOADER_ARGS diff --git a/docs/index.md b/docs/index.md index 90c37885e..0ad65288b 100644 --- a/docs/index.md +++ b/docs/index.md @@ -31,7 +31,7 @@ hide: _Oncology FM Evaluation Framework by [kaiko.ai](https://www.kaiko.ai/)_ -*eva* currently supports performance evaluation for vision Foundation Models ("FMs") and supervised machine learning models on WSI (patch- and slide-level), radiology image segmentation, and text classification tasks. +*eva* currently supports performance evaluation for Foundation Models ("FMs") accross multiple oncology domains and data modalities. With *eva* we provide the open-source community with an easy-to-use framework that follows industry best practices to deliver a robust, reproducible and fair evaluation benchmark across FMs of different sizes and architectures. @@ -69,6 +69,11 @@ Supported datasets & tasks include: - **[PubMedQA](datasets/pubmedqa.md)**: Medical question answering classification +*Multimodal datasets* + +- **[PatchCamelyon (image-text)](datasets/patch_camelyon.md)**: Vision-language benchmark variation for the popular vision Patch Camelyon task, where the goal is to classify breast cancer patches, using both the image and a text prompt. + + To evaluate FMs, *eva* provides support for different model-formats, including models trained with PyTorch, models available on HuggingFace and ONNX-models. For other formats custom wrappers can be implemented. diff --git a/docs/reference/language/models/networks.md b/docs/reference/language/models/networks.md index 56b26840d..fdf219726 100644 --- a/docs/reference/language/models/networks.md +++ b/docs/reference/language/models/networks.md @@ -2,4 +2,4 @@ Reference information for the language model `Networks` API. -::: eva.language.models.modules.TextModule \ No newline at end of file +::: eva.language.models.modules.LanguageModule \ No newline at end of file diff --git a/docs/reference/language/models/wrappers.md b/docs/reference/language/models/wrappers.md index dc39eab9b..3e9c7ec5c 100644 --- a/docs/reference/language/models/wrappers.md +++ b/docs/reference/language/models/wrappers.md @@ -2,6 +2,6 @@ Reference information for the language model `Wrappers` API. -::: eva.language.models.wrappers.HuggingFaceTextModel -::: eva.language.models.wrappers.LiteLLMTextModel -::: eva.language.models.wrappers.VLLMTextModel \ No newline at end of file +::: eva.language.models.wrappers.HuggingFaceModel +::: eva.language.models.wrappers.LiteLLMModel +::: eva.language.models.wrappers.VllmModel \ No newline at end of file diff --git a/docs/reference/multimodal/data/datasets.md b/docs/reference/multimodal/data/datasets.md new file mode 100644 index 000000000..da8837e22 --- /dev/null +++ b/docs/reference/multimodal/data/datasets.md @@ -0,0 +1,6 @@ +# Datasets + +Reference information for the multimodal data `Datasets` API. + +::: eva.multimodal.data.datasets.TextImageDataset +::: eva.multimodal.data.datasets.PatchCamelyon \ No newline at end of file diff --git a/docs/reference/multimodal/data/index.md b/docs/reference/multimodal/data/index.md new file mode 100644 index 000000000..fd7ac8b91 --- /dev/null +++ b/docs/reference/multimodal/data/index.md @@ -0,0 +1,3 @@ +# Data + +Multimodal data utilities for working with datasets that combine multiple modalities. \ No newline at end of file diff --git a/docs/reference/multimodal/index.md b/docs/reference/multimodal/index.md new file mode 100644 index 000000000..b2d37dc8f --- /dev/null +++ b/docs/reference/multimodal/index.md @@ -0,0 +1,8 @@ +# Multimodal + +Reference information for the `multimodal` API. + +If you have not already installed the `multimodal`-package, install it with: +``` +pip install 'kaiko-eva[multimodal]' +``` \ No newline at end of file diff --git a/docs/reference/multimodal/models/modules.md b/docs/reference/multimodal/models/modules.md new file mode 100644 index 000000000..da7f6cf56 --- /dev/null +++ b/docs/reference/multimodal/models/modules.md @@ -0,0 +1,5 @@ +# Modules + +Reference information for the multimodal `Modules` API. + +::: eva.multimodal.models.modules.VisionLanguageModule \ No newline at end of file diff --git a/docs/reference/multimodal/models/networks.md b/docs/reference/multimodal/models/networks.md new file mode 100644 index 000000000..59390b9ff --- /dev/null +++ b/docs/reference/multimodal/models/networks.md @@ -0,0 +1,14 @@ +# Networks + +Reference information for the multimodal `Networks` API. + +## Model Registry + +::: eva.multimodal.models.networks.model_registry + +## Pre-configured Models + +::: eva.multimodal.models.networks.Claude35Sonnet20240620 +::: eva.multimodal.models.networks.Claude37Sonnet20250219 +::: eva.multimodal.models.networks.PathoR13b +::: eva.multimodal.models.networks.Qwen25VL7BInstruct \ No newline at end of file diff --git a/docs/reference/multimodal/models/wrappers.md b/docs/reference/multimodal/models/wrappers.md new file mode 100644 index 000000000..bdb755249 --- /dev/null +++ b/docs/reference/multimodal/models/wrappers.md @@ -0,0 +1,8 @@ +# Wrappers + +Reference information for the multimodal `Wrappers` API. + +::: eva.multimodal.models.wrappers.VisionLanguageModel +::: eva.multimodal.models.wrappers.ModelFromRegistry +::: eva.multimodal.models.wrappers.HuggingFaceModel +::: eva.multimodal.models.wrappers.LiteLLMModel \ No newline at end of file diff --git a/docs/reference/multimodal/utils/image.md b/docs/reference/multimodal/utils/image.md new file mode 100644 index 000000000..277ae151b --- /dev/null +++ b/docs/reference/multimodal/utils/image.md @@ -0,0 +1,5 @@ +# Image Utilities + +Reference information for the multimodal image utilities API. + +::: eva.multimodal.utils.image.encode_image \ No newline at end of file diff --git a/docs/user-guide/tutorials/pubmedqa_classification.md b/docs/user-guide/tutorials/pubmedqa_classification.md index 09024828b..4c85db070 100644 --- a/docs/user-guide/tutorials/pubmedqa_classification.md +++ b/docs/user-guide/tutorials/pubmedqa_classification.md @@ -48,10 +48,10 @@ First, update the config to use the HuggingFace wrapper: ```yaml model: - class_path: eva.language.models.TextModule + class_path: eva.language.models.LanguageModule init_args: model: - class_path: eva.language.models.HuggingFaceTextModel + class_path: eva.language.models.HuggingFaceModel init_args: model_name_or_path: meta-llama/Llama-3.2-1B-Instruct ``` @@ -77,10 +77,10 @@ For larger models that require specialized infrastructure, you'll need to: ```yaml model: - class_path: eva.language.models.TextModule + class_path: eva.language.models.LanguageModule init_args: model: - class_path: eva.language.models.VLLMTextModel + class_path: eva.language.models.VllmModel init_args: model_name_or_path: meta-llama/Llama-2-70b-chat-hf ``` @@ -126,15 +126,10 @@ Once the evaluation is complete: The PubMedQA config demonstrates several important concepts: -#### Text prompting: -```yaml -prompt: "Instruction: You are an expert in biomedical research. Please carefully read the question and the relevant context and answer with yes, no, or maybe. Only answer with one of these three words." -``` - #### Model configuration (LiteLLM): ```yaml model: - class_path: eva.language.models.LiteLLMTextModel + class_path: eva.language.models.LiteLLMModel init_args: model_name_or_path: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-latest} ``` @@ -164,14 +159,6 @@ postprocess: ## Advanced usage -### Custom prompts - -You can experiment with different prompting strategies by modifying the prompt in the config file. For example, you might try: - -- Chain-of-thought prompting -- Few-shot examples -- Different output formats - ### Model comparison Run evaluations with multiple models to compare their performance on the 1000-question test set: diff --git a/mkdocs.yml b/mkdocs.yml index 71eaad52b..23102f3c3 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -152,3 +152,14 @@ nav: - Models: - Networks: reference/language/models/networks.md - Wrappers: reference/language/models/wrappers.md + - Multimodal: + - reference/multimodal/index.md + - Data: + - reference/multimodal/data/index.md + - reference/multimodal/data/datasets.md + - Models: + - Modules: reference/multimodal/models/modules.md + - Networks: reference/multimodal/models/networks.md + - Wrappers: reference/multimodal/models/wrappers.md + - Utils: + - Image: reference/multimodal/utils/image.md diff --git a/pdm.lock b/pdm.lock index e95cf3fcf..58cc650fb 100644 --- a/pdm.lock +++ b/pdm.lock @@ -5,7 +5,7 @@ groups = ["default", "all", "dev", "docs", "language", "lint", "test", "typecheck", "vision"] strategy = ["inherit_metadata"] lock_version = "4.5.0" -content_hash = "sha256:8b15139ebec0535e860fdedb3fdb86c4ee468fc8d661bcc9b34e7688d2673a2c" +content_hash = "sha256:13ce11cc3c2dead166d3465603d80e3fef059bfff06cef1e4e2c1deca7f78c46" [[metadata.targets]] requires_python = ">=3.10" @@ -207,6 +207,17 @@ files = [ {file = "babel-2.16.0.tar.gz", hash = "sha256:d1f3554ca26605fe173f3de0c65f750f5a42f924499bf134de6423582298e316"}, ] +[[package]] +name = "backoff" +version = "2.2.1" +requires_python = ">=3.7,<4.0" +summary = "Function decoration for backoff and retry" +groups = ["all", "language"] +files = [ + {file = "backoff-2.2.1-py3-none-any.whl", hash = "sha256:63579f9a0628e06278f7e47b7d7d5b6ce20dc65c5e96a6f3ca99a6adca0396e8"}, + {file = "backoff-2.2.1.tar.gz", hash = "sha256:03f829f5bb1923180821643f8753b0502c3b682293992485b0eef2807afa5cba"}, +] + [[package]] name = "backports-strenum" version = "1.3.1" diff --git a/pyproject.toml b/pyproject.toml index a0a679094..6e0a59c04 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,6 +75,11 @@ vision = [ language = [ "datasets<4.0.0,>=2.19.0", "litellm>=1.61.8", + "backoff>=2.2.1", +] +multimodal = [ + "litellm>=1.61.8", + "backoff>=2.2.1", ] all = [ "h5py>=3.10.0", @@ -91,6 +96,7 @@ all = [ "einops>=0.8.1", "datasets<4.0.0,>=2.19.0", "litellm>=1.61.8", + "backoff>=2.2.1", ] [project.scripts] diff --git a/src/eva/core/cli/setup.py b/src/eva/core/cli/setup.py index c45dd3002..fe5a48db3 100644 --- a/src/eva/core/cli/setup.py +++ b/src/eva/core/cli/setup.py @@ -59,7 +59,7 @@ def _initialize_logger() -> None: " :: {level}" " :: {message}", colorize=True, - level="INFO", + level=os.getenv("LOGURU_LEVEL", "INFO"), ) diff --git a/src/eva/core/data/dataloaders/__init__.py b/src/eva/core/data/dataloaders/__init__.py index 8bb9a8c67..f20954fc3 100644 --- a/src/eva/core/data/dataloaders/__init__.py +++ b/src/eva/core/data/dataloaders/__init__.py @@ -1,6 +1,5 @@ """Dataloaders API.""" -from eva.core.data.dataloaders.collate_fn import text_collate_fn from eva.core.data.dataloaders.dataloader import DataLoader -__all__ = ["text_collate_fn", "DataLoader"] +__all__ = ["DataLoader"] diff --git a/src/eva/core/data/dataloaders/collate_fn/__init__.py b/src/eva/core/data/dataloaders/collate_fn/__init__.py deleted file mode 100644 index e6421c9bc..000000000 --- a/src/eva/core/data/dataloaders/collate_fn/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Collate functions API.""" - -from eva.core.data.dataloaders.collate_fn.collate import text_collate_fn - -__all__ = ["text_collate_fn"] diff --git a/src/eva/core/data/dataloaders/collate_fn/collate.py b/src/eva/core/data/dataloaders/collate_fn/collate.py deleted file mode 100644 index e434a9b31..000000000 --- a/src/eva/core/data/dataloaders/collate_fn/collate.py +++ /dev/null @@ -1,24 +0,0 @@ -"""Collate functions for text data.""" - -from typing import Dict, List, Tuple - -import torch - - -def text_collate_fn( - batch: List[Tuple[str, torch.Tensor, Dict]], -) -> Tuple[List[str], torch.Tensor, List[Dict]]: - """Collate function for text data that keeps texts as separate strings. - - Args: - batch: List of tuples containing (text, target, metadata) from the dataset - - Returns: - Tuple containing: - - List of text strings - - Batched tensor of targets - - List of metadata dictionaries - """ - texts, targets, metadata = zip(*batch, strict=False) - targets = torch.stack(targets) - return list(texts), targets, list(metadata) diff --git a/src/eva/core/models/wrappers/base.py b/src/eva/core/models/wrappers/base.py index 28f13a06c..30e30861b 100644 --- a/src/eva/core/models/wrappers/base.py +++ b/src/eva/core/models/wrappers/base.py @@ -25,7 +25,7 @@ def __init__(self, transforms: Callable | None = None) -> None: self._output_transforms = transforms - self._model: Callable[..., OutputType] | nn.Module + self.model: Callable[..., OutputType] | nn.Module @override def forward(self, tensor: InputType) -> OutputType: @@ -43,7 +43,7 @@ def model_forward(self, tensor: InputType) -> OutputType: Args: tensor: The input tensor to the model. """ - return self._model(tensor) + return self.model(tensor) def _apply_transforms(self, tensor: OutputType) -> OutputType: if self._output_transforms is not None: diff --git a/src/eva/core/models/wrappers/from_function.py b/src/eva/core/models/wrappers/from_function.py index a968b5a5b..3cee623f8 100644 --- a/src/eva/core/models/wrappers/from_function.py +++ b/src/eva/core/models/wrappers/from_function.py @@ -41,12 +41,12 @@ def __init__( self._arguments = arguments self._checkpoint_path = checkpoint_path - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: + def load_model(self) -> nn.Module: class_path = jsonargparse.class_from_function(self._path, func_return=nn.Module) model = class_path(**self._arguments or {}) if self._checkpoint_path is not None: _utils.load_model_weights(model, self._checkpoint_path) - self._model = model + return model diff --git a/src/eva/core/models/wrappers/from_torchhub.py b/src/eva/core/models/wrappers/from_torchhub.py index 3b83bde1f..3757d21df 100644 --- a/src/eva/core/models/wrappers/from_torchhub.py +++ b/src/eva/core/models/wrappers/from_torchhub.py @@ -52,12 +52,12 @@ def __init__( self._trust_repo = trust_repo self._model_kwargs = model_kwargs or {} - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: + def load_model(self) -> nn.Module: """Builds and loads the torch.hub model.""" - self._model: nn.Module = torch.hub.load( + model: nn.Module = torch.hub.load( repo_or_dir=self._repo_or_dir, model=self._model_name, trust_repo=self._trust_repo, @@ -66,21 +66,23 @@ def load_model(self) -> None: ) # type: ignore if self._checkpoint_path: - _utils.load_model_weights(self._model, self._checkpoint_path) + _utils.load_model_weights(model, self._checkpoint_path) TorchHubModel.__name__ = self._model_name + return model + @override def model_forward(self, tensor: torch.Tensor) -> torch.Tensor | List[torch.Tensor]: if self._out_indices is not None: - if not hasattr(self._model, "get_intermediate_layers"): + if not hasattr(self.model, "get_intermediate_layers"): raise ValueError( "Only models with `get_intermediate_layers` are supported " "when using `out_indices`." ) return list( - self._model.get_intermediate_layers( + self.model.get_intermediate_layers( # type: ignore tensor, self._out_indices, reshape=True, @@ -89,4 +91,4 @@ def model_forward(self, tensor: torch.Tensor) -> torch.Tensor | List[torch.Tenso ) ) - return self._model(tensor) + return self.model(tensor) diff --git a/src/eva/core/models/wrappers/huggingface.py b/src/eva/core/models/wrappers/huggingface.py index 1d3b8abe6..950bc9b7b 100644 --- a/src/eva/core/models/wrappers/huggingface.py +++ b/src/eva/core/models/wrappers/huggingface.py @@ -4,6 +4,7 @@ import torch import transformers +from torch import nn from typing_extensions import override from eva.core.models.wrappers import base @@ -33,12 +34,10 @@ def __init__( self._model_name_or_path = model_name_or_path self._model_kwargs = model_kwargs or {} - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: + def load_model(self) -> nn.Module: # Use safetensors to avoid torch.load security vulnerability model_kwargs = {"use_safetensors": True, **self._model_kwargs} - self._model = transformers.AutoModel.from_pretrained( - self._model_name_or_path, **model_kwargs - ) + return transformers.AutoModel.from_pretrained(self._model_name_or_path, **model_kwargs) diff --git a/src/eva/core/models/wrappers/onnx.py b/src/eva/core/models/wrappers/onnx.py index 16046e7a1..f6b98fa8c 100644 --- a/src/eva/core/models/wrappers/onnx.py +++ b/src/eva/core/models/wrappers/onnx.py @@ -30,21 +30,21 @@ def __init__( self._path = path self._device = device - self.load_model() + self.model = self.load_model() @override def load_model(self) -> Any: if self._device == "cuda" and not torch.cuda.is_available(): raise ValueError("Device is set to 'cuda', but CUDA is not available.") provider = "CUDAExecutionProvider" if self._device == "cuda" else "CPUExecutionProvider" - self._model = ort.InferenceSession(self._path, providers=[provider]) # type: ignore + return ort.InferenceSession(self._path, providers=[provider]) # type: ignore @override def model_forward(self, tensor: torch.Tensor) -> torch.Tensor: # TODO: Use IO binding to avoid copying the tensor to CPU. # https://onnxruntime.ai/docs/api/python/api_summary.html#data-on-device - if not isinstance(self._model, ort.InferenceSession): + if not isinstance(self.model, ort.InferenceSession): raise ValueError("Model is not loaded.") - inputs = {self._model.get_inputs()[0].name: tensor.detach().cpu().numpy()} - outputs = self._model.run(None, inputs)[0] + inputs = {self.model.get_inputs()[0].name: tensor.detach().cpu().numpy()} + outputs = self.model.run(None, inputs)[0] return torch.from_numpy(outputs).float().to(tensor.device) diff --git a/src/eva/language/__init__.py b/src/eva/language/__init__.py index 8062734f0..1a37a87a2 100644 --- a/src/eva/language/__init__.py +++ b/src/eva/language/__init__.py @@ -1,6 +1,7 @@ """eva language API.""" try: + from eva.language import models from eva.language.data import datasets except ImportError as e: msg = ( @@ -10,4 +11,4 @@ ) raise ImportError(str(e) + "\n\n" + msg) from e -__all__ = ["datasets"] +__all__ = ["models", "datasets"] diff --git a/src/eva/language/data/dataloaders/__init__.py b/src/eva/language/data/dataloaders/__init__.py new file mode 100644 index 000000000..90763136f --- /dev/null +++ b/src/eva/language/data/dataloaders/__init__.py @@ -0,0 +1,5 @@ +"""Language Dataloaders API.""" + +from eva.language.data.dataloaders.collate_fn import text_collate + +__all__ = ["text_collate"] diff --git a/src/eva/language/data/dataloaders/collate_fn/__init__.py b/src/eva/language/data/dataloaders/collate_fn/__init__.py new file mode 100644 index 000000000..61671d160 --- /dev/null +++ b/src/eva/language/data/dataloaders/collate_fn/__init__.py @@ -0,0 +1,5 @@ +"""Collate functions API.""" + +from eva.language.data.dataloaders.collate_fn.text import text_collate + +__all__ = ["text_collate"] diff --git a/src/eva/language/data/dataloaders/collate_fn/text.py b/src/eva/language/data/dataloaders/collate_fn/text.py new file mode 100644 index 000000000..d69a188e5 --- /dev/null +++ b/src/eva/language/data/dataloaders/collate_fn/text.py @@ -0,0 +1,32 @@ +"""Collate functions for text data.""" + +from typing import List + +from torch.utils.data._utils.collate import default_collate + +from eva.language.data.datasets.typings import TextSample +from eva.language.models.typings import TextBatch + + +def text_collate(batch: List[TextSample]) -> TextBatch: + """Collate function for text data that keeps texts as separate strings. + + Args: + batch: List of tuples containing (text, target, metadata) from the dataset + + Returns: + A batch of text samples with targets and metadata. + """ + texts, targets, metadata = zip(*batch, strict=False) + first_sample = batch[0] + metadata = None + if first_sample.metadata is not None: + metadata = { + k: [sample.metadata[k] for sample in batch if sample.metadata] + for k in first_sample.metadata.keys() + } + return TextBatch( + text=list(texts), + target=default_collate(targets) if targets[0] is not None else None, + metadata=metadata, + ) diff --git a/src/eva/language/data/datasets/__init__.py b/src/eva/language/data/datasets/__init__.py index 171b0204f..e51806933 100644 --- a/src/eva/language/data/datasets/__init__.py +++ b/src/eva/language/data/datasets/__init__.py @@ -1,7 +1,7 @@ """Language Datasets API.""" +from eva.language.data.datasets.base import LanguageDataset from eva.language.data.datasets.classification import PubMedQA -from eva.language.data.datasets.language import LanguageDataset __all__ = [ "PubMedQA", diff --git a/src/eva/language/data/datasets/language.py b/src/eva/language/data/datasets/base.py similarity index 84% rename from src/eva/language/data/datasets/language.py rename to src/eva/language/data/datasets/base.py index 330690385..d07cc4b97 100644 --- a/src/eva/language/data/datasets/language.py +++ b/src/eva/language/data/datasets/base.py @@ -10,4 +10,4 @@ class LanguageDataset(base.MapDataset, abc.ABC, Generic[DataSample]): - """Base dataset class for text tasks.""" + """Base dataset class for language tasks.""" diff --git a/src/eva/language/data/datasets/classification/base.py b/src/eva/language/data/datasets/classification/base.py index c31699512..0dc043723 100644 --- a/src/eva/language/data/datasets/classification/base.py +++ b/src/eva/language/data/datasets/classification/base.py @@ -1,15 +1,13 @@ """Base for text classification datasets.""" -import abc -from typing import Any, Dict, List, Tuple +from typing import Dict, List import torch -from typing_extensions import override -from eva.language.data.datasets.language import LanguageDataset +from eva.language.data.datasets.text import TextDataset -class TextClassification(LanguageDataset[Tuple[str, torch.Tensor, Dict[str, Any]]], abc.ABC): +class TextClassification(TextDataset[torch.Tensor]): """Text classification abstract dataset.""" def __init__(self) -> None: @@ -23,41 +21,3 @@ def classes(self) -> List[str] | None: @property def class_to_idx(self) -> Dict[str, int] | None: """Returns class name to index mapping.""" - - def load_metadata(self, index: int) -> Dict[str, Any] | None: - """Returns the dataset metadata. - - Args: - index: The index of the data sample. - - Returns: - The sample metadata. - """ - - @abc.abstractmethod - def load_text(self, index: int) -> str: - """Returns the text content. - - Args: - index: The index of the data sample. - - Returns: - The text content. - """ - raise NotImplementedError - - @abc.abstractmethod - def load_target(self, index: int) -> torch.Tensor: - """Returns the target label. - - Args: - index: The index of the data sample. - - Returns: - The target label. - """ - raise NotImplementedError - - @override - def __getitem__(self, index: int) -> Tuple[str, torch.Tensor, Dict[str, Any]]: - return (self.load_text(index), self.load_target(index), self.load_metadata(index) or {}) diff --git a/src/eva/language/data/datasets/classification/pubmedqa.py b/src/eva/language/data/datasets/classification/pubmedqa.py index 5673321c5..cd09bf88b 100644 --- a/src/eva/language/data/datasets/classification/pubmedqa.py +++ b/src/eva/language/data/datasets/classification/pubmedqa.py @@ -10,6 +10,7 @@ from typing_extensions import override from eva.language.data.datasets.classification import base +from eva.language.data.messages import MessageSeries, UserMessage class PubMedQA(base.TextClassification): @@ -114,11 +115,18 @@ def class_to_idx(self) -> Dict[str, int]: return {"no": 0, "yes": 1, "maybe": 2} @override - def load_text(self, index: int) -> str: + def load_text(self, index: int) -> MessageSeries: if index < 0 or index >= len(self.dataset): raise IndexError(f"Index {index} out of range for dataset of size {len(self.dataset)}") sample = dict(self.dataset[index]) - return f"Question: {sample['QUESTION']}\nContext: " + " ".join(sample["CONTEXTS"]) + return [ + UserMessage( + content=f"Question: {sample['QUESTION']}\nContext: " + + " ".join(sample["CONTEXTS"]) + + "\nInstruction: Carefully read the question and the provided context. " + + "Answer with one word: 'yes', 'no', or 'maybe'. Answer: " + ) + ] @override def load_target(self, index: int) -> torch.Tensor: diff --git a/src/eva/language/data/datasets/schemas.py b/src/eva/language/data/datasets/schemas.py new file mode 100644 index 000000000..02359614e --- /dev/null +++ b/src/eva/language/data/datasets/schemas.py @@ -0,0 +1,15 @@ +"""Schema definitions for dataset classes.""" + +import dataclasses +from typing import Callable + + +@dataclasses.dataclass(frozen=True) +class TransformsSchema: + """Schema for dataset transforms.""" + + text: Callable | None = None + """Text transformation""" + + target: Callable | None = None + """Target transformation""" diff --git a/src/eva/language/data/datasets/text.py b/src/eva/language/data/datasets/text.py new file mode 100644 index 000000000..f2e3a71af --- /dev/null +++ b/src/eva/language/data/datasets/text.py @@ -0,0 +1,93 @@ +"""Base classes for text-image datasets.""" + +import abc +from typing import Any, Dict, Generic + +from typing_extensions import override + +from eva.language.data.datasets.base import LanguageDataset +from eva.language.data.datasets.schemas import TransformsSchema +from eva.language.data.datasets.typings import TargetType, TextSample +from eva.language.data.messages import MessageSeries + + +class TextDataset(LanguageDataset[TextSample[TargetType]], abc.ABC, Generic[TargetType]): + """Base dataset class for text-based tasks.""" + + def __init__(self, *args, transforms: TransformsSchema | None = None, **kwargs) -> None: + """Initializes the dataset. + + Args: + *args: Positional arguments for the base class. + transforms: The transforms to apply to the text and target when + loading the samples. + **kwargs: Keyword arguments for the base class. + """ + super().__init__(*args, **kwargs) + + self.transforms = transforms + + def load_metadata(self, index: int) -> Dict[str, Any] | None: + """Returns the dataset metadata. + + Args: + index: The index of the data sample. + + Returns: + The sample metadata. + """ + + @abc.abstractmethod + def load_text(self, index: int) -> MessageSeries: + """Returns the text content. + + Args: + index: The index of the data sample. + + Returns: + The text content. + """ + raise NotImplementedError + + @abc.abstractmethod + def load_target(self, index: int) -> TargetType: + """Returns the target label. + + Args: + index: The index of the data sample. + + Returns: + The target label. + """ + raise NotImplementedError + + @override + def __getitem__(self, index: int) -> TextSample[TargetType]: + item = TextSample( + text=self.load_text(index), + target=self.load_target(index), + metadata=self.load_metadata(index) or {}, + ) + return self._apply_transforms(item) + + def _apply_transforms(self, sample: TextSample[TargetType]) -> TextSample[TargetType]: + """Applies the dataset transforms to the text and target. + + Args: + sample: The text sample.. + target: The target label. + + Returns: + The transformed sample. + """ + if self.transforms: + text = self.transforms.text(sample.text) if self.transforms.text else sample.text + target = ( + self.transforms.target(sample.target) if self.transforms.target else sample.target + ) + return TextSample( + text=text, + target=target, + metadata=sample.metadata, + ) + return sample diff --git a/src/eva/language/data/datasets/typings.py b/src/eva/language/data/datasets/typings.py new file mode 100644 index 000000000..0dd5b1b1f --- /dev/null +++ b/src/eva/language/data/datasets/typings.py @@ -0,0 +1,23 @@ +"""Typings for multimodal datasets.""" + +from typing import Any, Generic, TypeVar + +from typing_extensions import NamedTuple + +from eva.language.data.messages import MessageSeries + +TargetType = TypeVar("TargetType") +"""The target data type.""" + + +class TextSample(NamedTuple, Generic[TargetType]): + """Text sample with target and metadata.""" + + text: MessageSeries + """One or multiple conversation messages.""" + + target: TargetType | None + """Target data.""" + + metadata: dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/language/data/messages.py b/src/eva/language/data/messages.py new file mode 100644 index 000000000..007d37797 --- /dev/null +++ b/src/eva/language/data/messages.py @@ -0,0 +1,51 @@ +"""Types and classes for conversation messages in a multimodal context.""" + +import dataclasses +from typing import Any, Dict, List + + +@dataclasses.dataclass +class Message: + """Base class for a message in a conversation.""" + + content: str + role: str + + def to_dict(self) -> Dict[str, Any]: + """Convert the message to a dictionary.""" + return dataclasses.asdict(self) + + +@dataclasses.dataclass +class UserMessage(Message): + """User message in a conversation.""" + + role: str = "user" + + +@dataclasses.dataclass +class AssistantMessage(Message): + """Assistant message in a conversation.""" + + role: str = "assistant" + + +@dataclasses.dataclass +class SystemMessage(Message): + """System message in a conversation.""" + + role: str = "system" + + +@dataclasses.dataclass +class ModelSystemMessage(SystemMessage): + """System message for model-specific instructions.""" + + +@dataclasses.dataclass +class TaskSystemMessage(SystemMessage): + """System message for task-specific instructions.""" + + +MessageSeries = List[Message] +"""A series of conversation messages, can contain a mix of system, user, and AI messages.""" diff --git a/src/eva/language/models/__init__.py b/src/eva/language/models/__init__.py index 043e542b8..c8d3bf192 100644 --- a/src/eva/language/models/__init__.py +++ b/src/eva/language/models/__init__.py @@ -1,25 +1,27 @@ """Language Models API.""" -from eva.language.models import modules, wrappers -from eva.language.models.modules import TextModule -from eva.language.models.wrappers import HuggingFaceTextModel, LiteLLMTextModel +from eva.language.models import modules, networks, wrappers +from eva.language.models.modules import LanguageModule +from eva.language.models.wrappers import HuggingFaceModel, LiteLLMModel try: - from eva.language.models.wrappers import VLLMTextModel + from eva.language.models.wrappers import VllmModel __all__ = [ "modules", "wrappers", - "TextModule", - "HuggingFaceTextModel", - "LiteLLMTextModel", - "VLLMTextModel", + "networks", + "HuggingFaceModel", + "LiteLLMModel", + "VllmModel", + "LanguageModule", ] except ImportError: __all__ = [ "modules", "wrappers", - "TextModule", - "HuggingFaceTextModel", - "LiteLLMTextModel", + "networks", + "HuggingFaceModel", + "LiteLLMModel", + "LanguageModule", ] diff --git a/src/eva/language/models/modules/__init__.py b/src/eva/language/models/modules/__init__.py index e770290be..3dcd2cfc2 100644 --- a/src/eva/language/models/modules/__init__.py +++ b/src/eva/language/models/modules/__init__.py @@ -1,5 +1,5 @@ """Language Networks API.""" -from eva.language.models.modules.text import TextModule +from eva.language.models.modules.language import LanguageModule -__all__ = ["TextModule"] +__all__ = ["LanguageModule"] diff --git a/src/eva/language/models/modules/language.py b/src/eva/language/models/modules/language.py new file mode 100644 index 000000000..74cb3ef87 --- /dev/null +++ b/src/eva/language/models/modules/language.py @@ -0,0 +1,55 @@ +"""Model module for language models.""" + +from typing import Any, List + +from lightning.pytorch.utilities.types import STEP_OUTPUT +from torch import nn +from typing_extensions import override + +from eva.core.metrics import structs as metrics_lib +from eva.core.models.modules import module +from eva.core.models.modules.utils import batch_postprocess +from eva.language.models.typings import TextBatch + + +class LanguageModule(module.ModelModule): + """Model module for language tasks.""" + + def __init__( + self, + model: nn.Module, + metrics: metrics_lib.MetricsSchema | None = None, + postprocess: batch_postprocess.BatchPostProcess | None = None, + ) -> None: + """Initializes the text inference module. + + Args: + model: Model instance to use for forward pass. + metrics: Metrics schema for evaluation. + postprocess: A helper function to post-process model outputs before evaluation. + """ + super().__init__(metrics=metrics, postprocess=postprocess) + + self.model = model + + @override + def forward(self, batch: TextBatch, *args: Any, **kwargs: Any) -> List[str]: + return self.model(batch) + + @override + def validation_step(self, batch: TextBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + @override + def test_step(self, batch: TextBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + def _batch_step(self, batch: TextBatch) -> STEP_OUTPUT: + text, targets, metadata = TextBatch(*batch) + predictions = self.forward(batch) + return { + "inputs": text, + "predictions": predictions, + "targets": targets, + "metadata": metadata, + } diff --git a/src/eva/language/models/modules/text.py b/src/eva/language/models/modules/text.py deleted file mode 100644 index 5f561328e..000000000 --- a/src/eva/language/models/modules/text.py +++ /dev/null @@ -1,85 +0,0 @@ -"""LLM Text Module for Inference.""" - -from typing import Any, List - -from lightning.pytorch.utilities.types import STEP_OUTPUT -from loguru import logger -from torch import nn -from typing_extensions import override - -from eva.core.metrics import structs as metrics_lib -from eva.core.models.modules import module -from eva.core.models.modules.utils import batch_postprocess -from eva.language.models.modules.typings import TEXT_BATCH - - -class TextModule(module.ModelModule): - """Text-based LLM module for inference. - - Uses LLM wrappers for text generation and supports evaluation using - configurable metrics and post-processing transforms. - """ - - def __init__( - self, - model: nn.Module, - prompt: str, - metrics: metrics_lib.MetricsSchema | None = None, - postprocess: batch_postprocess.BatchPostProcess | None = None, - ) -> None: - """Initializes the text inference module. - - Args: - model: An LLM wrapper (PyTorch-compatible) for text generation. - prompt: The prompt to use for generating text. - metrics: Metrics schema for evaluation. - postprocess: A helper function to post-process model outputs before evaluation. - """ - super().__init__(metrics=metrics, postprocess=postprocess) - - self.model = model - self.prompt = prompt - - @override - def forward(self, prompts: List[str], *args: Any, **kwargs: Any) -> List[str]: - """Generates text responses for a batch of prompts. - - Args: - prompts: List of input texts to generate responses. - args: Additional arguments. - kwargs: Additional keyword arguments. - - Returns: - List of generated responses. - """ - return self.model(prompts) - - @override - def validation_step(self, batch: TEXT_BATCH, *args: Any, **kwargs: Any) -> STEP_OUTPUT: - """Validation step that runs batch inference and evaluates metrics. - - Args: - batch: An input batch. - args: Additional arguments. - kwargs: Additional keyword arguments. - - Returns: - Dictionary with predictions, ground truth, and evaluation metrics. - """ - return self._batch_step(batch) - - def _batch_step(self, batch: TEXT_BATCH) -> STEP_OUTPUT: - """Runs inference on a batch and evaluates model predictions. - - Args: - batch: Input batch containing data, targets, and metadata. - - Returns: - Dictionary with predictions, ground truth, and evaluation metrics. - """ - data, targets, metadata = batch - messages = [str(d) + "\n" + self.prompt for d in data] - predictions = self(messages) - logger.debug(f"Predictions: {predictions}") - logger.debug(f"Targets: {targets}") - return {"predictions": predictions, "targets": targets, "metadata": metadata} diff --git a/src/eva/language/models/networks/__init__.py b/src/eva/language/models/networks/__init__.py new file mode 100644 index 000000000..707b3fcd1 --- /dev/null +++ b/src/eva/language/models/networks/__init__.py @@ -0,0 +1,12 @@ +"""Language networks API.""" + +from eva.language.models.networks.alibaba import Qwen205BInstruct +from eva.language.models.networks.api import Claude35Sonnet20240620, Claude37Sonnet20250219 +from eva.language.models.networks.registry import model_registry + +__all__ = [ + "Claude35Sonnet20240620", + "Claude37Sonnet20250219", + "Qwen205BInstruct", + "model_registry", +] diff --git a/src/eva/language/models/networks/alibaba.py b/src/eva/language/models/networks/alibaba.py new file mode 100644 index 000000000..4f25b9c88 --- /dev/null +++ b/src/eva/language/models/networks/alibaba.py @@ -0,0 +1,26 @@ +"""Models from Alibaba.""" + +import torch + +from eva.language.models import wrappers +from eva.language.models.networks.registry import model_registry + + +@model_registry.register("alibaba/qwen2-0-5b-instruct") +class Qwen205BInstruct(wrappers.HuggingFaceModel): + """Qwen2 0.5B Instruct model.""" + + def __init__(self, system_prompt: str | None = None, cache_dir: str | None = None): + """Initialize the model.""" + super().__init__( + model_name_or_path="Qwen/Qwen2-0.5B-Instruct", + model_kwargs={ + "torch_dtype": torch.bfloat16, + "cache_dir": cache_dir, + }, + generation_kwargs={ + "max_new_tokens": 512, + }, + system_prompt=system_prompt, + chat_mode=True, + ) diff --git a/src/eva/language/models/networks/api/__init__.py b/src/eva/language/models/networks/api/__init__.py new file mode 100644 index 000000000..10d777423 --- /dev/null +++ b/src/eva/language/models/networks/api/__init__.py @@ -0,0 +1,11 @@ +"""Multimodal API networks.""" + +from eva.language.models.networks.api.anthropic import ( + Claude35Sonnet20240620, + Claude37Sonnet20250219, +) + +__all__ = [ + "Claude35Sonnet20240620", + "Claude37Sonnet20250219", +] diff --git a/src/eva/language/models/networks/api/anthropic.py b/src/eva/language/models/networks/api/anthropic.py new file mode 100644 index 000000000..3c4a1522f --- /dev/null +++ b/src/eva/language/models/networks/api/anthropic.py @@ -0,0 +1,34 @@ +"""Models from Anthropic.""" + +import os + +from eva.language.models import wrappers +from eva.language.models.networks.registry import model_registry + + +class _Claude(wrappers.LiteLLMModel): + """Base class for Claude models.""" + + def __init__(self, model_name: str, system_prompt: str | None = None): + if not os.getenv("ANTHROPIC_API_KEY"): + raise ValueError("ANTHROPIC_API_KEY env variable must be set.") + + super().__init__(model_name=model_name, system_prompt=system_prompt) + + +@model_registry.register("anthropic/claude-3-5-sonnet-20240620") +class Claude35Sonnet20240620(_Claude): + """Claude 3.5 Sonnet (June 2024) model.""" + + def __init__(self, system_prompt: str | None = None): + """Initialize the model.""" + super().__init__(model_name="claude-3-5-sonnet-20240620", system_prompt=system_prompt) + + +@model_registry.register("anthropic/claude-3-7-sonnet-20250219") +class Claude37Sonnet20250219(_Claude): + """Claude 3.7 Sonnet (February 2025) model.""" + + def __init__(self, system_prompt: str | None = None): + """Initialize the model.""" + super().__init__(model_name="claude-3-7-sonnet-20250219", system_prompt=system_prompt) diff --git a/src/eva/language/models/networks/registry.py b/src/eva/language/models/networks/registry.py new file mode 100644 index 000000000..c65167c04 --- /dev/null +++ b/src/eva/language/models/networks/registry.py @@ -0,0 +1,5 @@ +"""Language Model Registry.""" + +from eva.core.utils.registry import Registry + +model_registry = Registry() diff --git a/src/eva/language/models/typings.py b/src/eva/language/models/typings.py new file mode 100644 index 000000000..71b35f019 --- /dev/null +++ b/src/eva/language/models/typings.py @@ -0,0 +1,23 @@ +"""Type definitions for language models.""" + +from typing import Any, Dict, Generic, List, TypeVar + +from typing_extensions import NamedTuple + +from eva.language.data.messages import MessageSeries + +TargetType = TypeVar("TargetType") +"""The target data type.""" + + +class TextBatch(NamedTuple, Generic[TargetType]): + """Text sample with target and metadata.""" + + text: List[MessageSeries] + """Text content.""" + + target: TargetType | None + """Target data.""" + + metadata: Dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/language/models/wrappers/__init__.py b/src/eva/language/models/wrappers/__init__.py index c36af634a..00482d826 100644 --- a/src/eva/language/models/wrappers/__init__.py +++ b/src/eva/language/models/wrappers/__init__.py @@ -1,11 +1,12 @@ """Language Model Wrappers API.""" -from eva.language.models.wrappers.huggingface import HuggingFaceTextModel -from eva.language.models.wrappers.litellm import LiteLLMTextModel +from eva.language.models.wrappers.from_registry import ModelFromRegistry +from eva.language.models.wrappers.huggingface import HuggingFaceModel +from eva.language.models.wrappers.litellm import LiteLLMModel try: - from eva.language.models.wrappers.vllm import VLLMTextModel + from eva.language.models.wrappers.vllm import VllmModel - __all__ = ["HuggingFaceTextModel", "LiteLLMTextModel", "VLLMTextModel"] + __all__ = ["HuggingFaceModel", "LiteLLMModel", "VllmModel", "ModelFromRegistry"] except ImportError: - __all__ = ["HuggingFaceTextModel", "LiteLLMTextModel"] + __all__ = ["HuggingFaceModel", "LiteLLMModel", "ModelFromRegistry"] diff --git a/src/eva/language/models/wrappers/base.py b/src/eva/language/models/wrappers/base.py new file mode 100644 index 000000000..548586255 --- /dev/null +++ b/src/eva/language/models/wrappers/base.py @@ -0,0 +1,47 @@ +"""Base class for language model wrappers.""" + +import abc +from typing import Any, Callable, List + +from typing_extensions import override + +from eva.core.models.wrappers import base +from eva.language.data.messages import ModelSystemMessage +from eva.language.models.typings import TextBatch + + +class LanguageModel(base.BaseModel[TextBatch, List[str]]): + """Base class for language models. + + Classes that inherit from this should implement the following methods: + - `load_model`: Loads & instantiates the model. + - `model_forward`: Implements the forward pass of the model. For API models, + this can be an API call. + - `format_inputs`: Preprocesses and converts the input batch into the format + expected by the `model_forward` method. + """ + + def __init__( + self, system_prompt: str | None, output_transforms: Callable | None = None + ) -> None: + """Creates a new model instance. + + Args: + system_prompt: The system prompt to use for the model (optional). + output_transforms: Optional transforms to apply to the output of + the model's forward pass. + """ + super().__init__(transforms=output_transforms) + + self.system_message = ModelSystemMessage(content=system_prompt) if system_prompt else None + + @override + def forward(self, batch: TextBatch) -> List[str]: + """Forward pass of the model.""" + inputs = self.format_inputs(batch) + return super().forward(inputs) + + @abc.abstractmethod + def format_inputs(self, batch: TextBatch) -> Any: + """Converts the inputs into the format expected by the model.""" + raise NotImplementedError diff --git a/src/eva/language/models/wrappers/from_registry.py b/src/eva/language/models/wrappers/from_registry.py new file mode 100644 index 000000000..560651fa9 --- /dev/null +++ b/src/eva/language/models/wrappers/from_registry.py @@ -0,0 +1,54 @@ +"""Vision backbone helper class.""" + +from typing import Any, Callable, Dict, List + +from torch import nn +from typing_extensions import override + +from eva.core.models.wrappers import base +from eva.core.utils import factory +from eva.language.models.networks.registry import model_registry +from eva.language.models.typings import TextBatch + + +class ModelFromRegistry(base.BaseModel[TextBatch, List[str]]): + """Wrapper class for vision backbone models. + + This class can be used by load backbones available in eva's + model registry by name. New backbones can be registered by using + the `@backbone_registry.register(model_name)` decorator. + """ + + def __init__( + self, + model_name: str, + model_kwargs: Dict[str, Any] | None = None, + model_extra_kwargs: Dict[str, Any] | None = None, + transforms: Callable | None = None, + ) -> None: + """Initializes the model. + + Args: + model_name: The name of the model to load. + model_kwargs: The arguments used for instantiating the model. + model_extra_kwargs: Extra arguments used for instantiating the model. + transforms: The transforms to apply to the output tensor + produced by the model. + """ + super().__init__(transforms=transforms) + + self._model_name = model_name + self._model_kwargs = model_kwargs or {} + self._model_extra_kwargs = model_extra_kwargs or {} + + self.model = self.load_model() + + @override + def load_model(self) -> nn.Module: + ModelFromRegistry.__name__ = self._model_name + + return factory.ModuleFactory( + registry=model_registry, + name=self._model_name, + init_args=self._model_kwargs | self._model_extra_kwargs, + ) diff --git a/src/eva/language/models/wrappers/huggingface.py b/src/eva/language/models/wrappers/huggingface.py index 0ceece62d..22a49cf2f 100644 --- a/src/eva/language/models/wrappers/huggingface.py +++ b/src/eva/language/models/wrappers/huggingface.py @@ -1,14 +1,16 @@ """LLM wrapper for HuggingFace `transformers` models.""" -from typing import Any, Dict, List, Literal +from typing import Any, Callable, Dict, List, Literal from transformers.pipelines import pipeline from typing_extensions import override -from eva.core.models.wrappers import base +from eva.language.models.typings import TextBatch +from eva.language.models.wrappers import base +from eva.language.utils.text import messages as message_utils -class HuggingFaceTextModel(base.BaseModel[List[str], List[str]]): +class HuggingFaceModel(base.LanguageModel): """Wrapper class for loading HuggingFace `transformers` models using pipelines.""" def __init__( @@ -16,7 +18,9 @@ def __init__( model_name_or_path: str, task: Literal["text-generation"] = "text-generation", model_kwargs: Dict[str, Any] | None = None, + system_prompt: str | None = None, generation_kwargs: Dict[str, Any] | None = None, + chat_mode: bool = True, ) -> None: """Initializes the model. @@ -26,27 +30,59 @@ def __init__( model hub. task: The pipeline task. Defaults to "text-generation". model_kwargs: Additional arguments for configuring the pipeline. + system_prompt: System prompt to use. generation_kwargs: Additional generation parameters (temperature, max_length, etc.). + chat_mode: Whether the specified model expects chat style messages. If set to False + the model is assumed to be a standard text completion model and will expect + plain text string inputs. """ - super().__init__() + super().__init__(system_prompt=system_prompt) self._model_name_or_path = model_name_or_path self._task = task self._model_kwargs = model_kwargs or {} self._generation_kwargs = generation_kwargs or {} + self._chat_mode = chat_mode - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: + def load_model(self) -> Callable: """Loads the model as a Hugging Face pipeline.""" - self._pipeline = pipeline( + return pipeline( task=self._task, model=self._model_name_or_path, trust_remote_code=True, **self._model_kwargs, ) + @override + def format_inputs(self, batch: TextBatch) -> List[List[Dict[str, Any]]] | List[str]: + """Formats inputs for HuggingFace models. + + Note: If multiple system messages are present, they will be combined + into a single message, given that many models only support a single + system prompt. + + Args: + batch: A batch of text and image inputs. + + Returns: + When in chat mode, returns a batch of message series following + OpenAI's API format {"role": "user", "content": "..."}, for non-chat + models returns a list of plain text strings. + """ + message_batch, _, _ = TextBatch(*batch) + message_batch = message_utils.batch_insert_system_message( + message_batch, self.system_message + ) + message_batch = list(map(message_utils.combine_system_messages, message_batch)) + + if self._chat_mode: + return list(map(message_utils.format_chat_message, message_batch)) + else: + return list(map(message_utils.merge_message_contents, message_batch)) + @override def model_forward(self, prompts: List[str]) -> List[str]: """Generates text using the pipeline. @@ -57,7 +93,7 @@ def model_forward(self, prompts: List[str]) -> List[str]: Returns: The generated text as a string. """ - outputs = self._pipeline(prompts, return_full_text=False, **self._generation_kwargs) + outputs = self.model(prompts, return_full_text=False, **self._generation_kwargs) if outputs is None: raise ValueError("Outputs from the model are None.") results = [] diff --git a/src/eva/language/models/wrappers/litellm.py b/src/eva/language/models/wrappers/litellm.py index 7c8093c54..751294e38 100644 --- a/src/eva/language/models/wrappers/litellm.py +++ b/src/eva/language/models/wrappers/litellm.py @@ -1,77 +1,112 @@ -"""LLM wrapper for litellm models.""" +"""LiteLLM language model wrapper.""" +import logging from typing import Any, Dict, List -from litellm import batch_completion # type: ignore +import backoff +import litellm +from litellm import batch_completion +from litellm.exceptions import ( + APIConnectionError, + InternalServerError, + RateLimitError, + ServiceUnavailableError, + Timeout, +) from loguru import logger from typing_extensions import override -from eva.core.models.wrappers import base +from eva.language.models.typings import TextBatch +from eva.language.models.wrappers import base +from eva.language.utils.text import messages as message_utils +RETRYABLE_ERRORS = ( + RateLimitError, + Timeout, + InternalServerError, + APIConnectionError, + ServiceUnavailableError, +) -class LiteLLMTextModel(base.BaseModel[List[str], List[str]]): - """Wrapper class for using litellm for chat-based text generation. - This wrapper uses litellm's `completion` function which accepts a list of - message dicts. The `forward` method converts a string prompt into a chat - message with a default "user" role, optionally prepends a system message, - and includes an API key if provided. - """ +class LiteLLMModel(base.LanguageModel): + """Wrapper class for LiteLLM language models.""" def __init__( self, - model_name_or_path: str, + model_name: str, model_kwargs: Dict[str, Any] | None = None, - ) -> None: - """Initializes the litellm chat model wrapper. + system_prompt: str | None = None, + log_level: int | None = logging.INFO, + ): + """Initialize the LiteLLM Wrapper. Args: - model_name_or_path: The model identifier (or name) for litellm - (e.g.,"openai/gpt-4o" or "anthropic/claude-3-sonnet-20240229"). + model_name: The name of the model to use. model_kwargs: Additional keyword arguments to pass during generation (e.g., `temperature`, `max_tokens`). + system_prompt: The system prompt to use (optional). + log_level: Optional logging level for LiteLLM. Defaults to WARNING. """ - super().__init__() - self._model_name_or_path = model_name_or_path - self._model_kwargs = model_kwargs or {} - self.load_model() + super().__init__(system_prompt=system_prompt) - @override - def load_model(self) -> None: - """Prepares the litellm model. + self.model_name = model_name + self.model_kwargs = model_kwargs or {} - Note: - litellm doesn't require an explicit loading step; models are called - directly during generation. This method exists for API consistency. - """ - pass + litellm.suppress_debug_info = True + + if log_level is not None: + logging.getLogger("LiteLLM").setLevel(log_level) @override - def model_forward(self, prompts: List[str]) -> List[str]: - """Generates text using litellm. + def format_inputs(self, batch: TextBatch) -> List[List[Dict[str, Any]]]: + """Formats inputs for LiteLLM. Args: - prompts: A list of prompts to be converted into a "user" message. + batch: A batch of text inputs. Returns: - A list of generated text responses. Failed generations will contain - error messages instead of generated text. + A list of messages in the following format: + [ + { + "role": ... + "content": ... + }, + ... + ] """ - messages = [[{"role": "user", "content": prompt}] for prompt in prompts] + message_batch, _, _ = TextBatch(*batch) - responses = batch_completion( - model=self._model_name_or_path, - messages=messages, - **self._model_kwargs, + message_batch = message_utils.batch_insert_system_message( + message_batch, self.system_message ) + message_batch = list(map(message_utils.combine_system_messages, message_batch)) - results = [] - for i, response in enumerate(responses): - if isinstance(response, Exception): - error_msg = f"Error generating text for prompt {i}: {response}" - logger.error(error_msg) - raise RuntimeError(error_msg) - else: - results.append(response["choices"][0]["message"]["content"]) + return list(map(message_utils.format_chat_message, message_batch)) - return results + @override + @backoff.on_exception( + backoff.expo, + RETRYABLE_ERRORS, + max_tries=20, + jitter=backoff.full_jitter, + on_backoff=lambda details: logger.warning( + f"Retrying due to {details.get('exception') or 'Unknown error'}" + ), + ) + def model_forward(self, batch: List[List[Dict[str, Any]]]) -> List[str]: + """Generates output text through API calls via LiteLLM's batch completion functionality.""" + outputs = batch_completion(model=self.model_name, messages=batch, **self.model_kwargs) + self._raise_exceptions(outputs) + + return [ + output["choices"][0]["message"]["content"] + for output in outputs + if output["choices"][0]["message"]["role"] == "assistant" + ] + + def _raise_exceptions(self, outputs: list): + for output in outputs: + if isinstance(output, Exception): + logger.error(f"Model {self.model_name} encountered an error: {output}") + raise output diff --git a/src/eva/language/models/wrappers/vllm.py b/src/eva/language/models/wrappers/vllm.py index 1b0d55142..60c6dbca4 100644 --- a/src/eva/language/models/wrappers/vllm.py +++ b/src/eva/language/models/wrappers/vllm.py @@ -1,6 +1,6 @@ """LLM wrapper for vLLM models.""" -from typing import Any, Dict, List, Sequence +from typing import Any, Dict, List from loguru import logger from typing_extensions import override @@ -11,17 +11,20 @@ from vllm.transformers_utils.tokenizer import AnyTokenizer # type: ignore except ImportError as e: raise ImportError( - "vLLM is required for VLLMTextModel but not installed. " + "vLLM is required for VllmModel but not installed. " "vLLM must be installed manually as it requires CUDA and is not included in dependencies. " "Install with: pip install vllm " "Note: vLLM requires Linux with CUDA support for optimal performance. " - "For alternatives, consider using HuggingFaceTextModel or LiteLLMTextModel." + "For alternatives, consider using HuggingFaceModel or LiteLLMModel." ) from e -from eva.core.models.wrappers import base +from eva.language.data.messages import MessageSeries +from eva.language.models.typings import TextBatch +from eva.language.models.wrappers import base +from eva.language.utils.text import messages as message_utils -class VLLMTextModel(base.BaseModel): +class VllmModel(base.LanguageModel): """Wrapper class for using vLLM for text generation. This wrapper loads a vLLM model, sets up the tokenizer and sampling @@ -34,6 +37,7 @@ def __init__( self, model_name_or_path: str, model_kwargs: Dict[str, Any] | None = None, + system_prompt: str | None = None, generation_kwargs: Dict[str, Any] | None = None, ) -> None: """Initializes the vLLM model wrapper. @@ -44,12 +48,13 @@ def __init__( model_kwargs: Arguments required to initialize the vLLM model, see [link](https://github.com/vllm-project/vllm/blob/main/vllm/entrypoints/llm.py) for more information. + system_prompt: System prompt to use. generation_kwargs: Arguments required to generate the output, need to align with the arguments of [vllm.SamplingParams](https://github.com/vllm-project/vllm/blob/main/vllm/sampling_params.py). """ - super().__init__() + super().__init__(system_prompt=system_prompt) self._model_name_or_path = model_name_or_path self._model_kwargs = model_kwargs or {} self._generation_kwargs = generation_kwargs or {} @@ -71,11 +76,11 @@ def load_model(self) -> None: raise RuntimeError("Model not initialized") self._llm_tokenizer = self._llm_model.get_tokenizer() - def _apply_chat_template(self, prompts: Sequence[str]) -> list[TokensPrompt]: + def _tokenize_messages(self, messages: List[MessageSeries]) -> List[TokensPrompt]: """Apply chat template to the messages. Args: - prompts: List of raw user strings. + messages: List of raw user strings. Returns: List of encoded messages. @@ -90,7 +95,8 @@ def _apply_chat_template(self, prompts: Sequence[str]) -> list[TokensPrompt]: if not hasattr(self._llm_tokenizer, "chat_template"): raise ValueError("Tokenizer does not have a chat template.") - chat_messages = [[{"role": "user", "content": p}] for p in prompts] + chat_messages = list(map(message_utils.format_chat_message, messages)) + encoded_messages = self._llm_tokenizer.apply_chat_template( chat_messages, # type: ignore tokenize=True, @@ -131,11 +137,30 @@ def _apply_chat_template(self, prompts: Sequence[str]) -> list[TokensPrompt]: return result - def generate(self, prompts: List[str]) -> List[str]: + @override + def format_inputs(self, batch: TextBatch) -> List[TokensPrompt]: + """Formats inputs for vLLM models. + + Args: + batch: A batch of text and image inputs. + + Returns: + List of formatted prompts. + """ + message_batch, _, _ = TextBatch(*batch) + message_batch = message_utils.batch_insert_system_message( + message_batch, self.system_message + ) + message_batch = list(map(message_utils.combine_system_messages, message_batch)) + + return self._tokenize_messages(message_batch) + + @override + def model_forward(self, batch: List[TokensPrompt]) -> List[str]: """Generates text for the given prompt using the vLLM model. Args: - prompts: A list of string prompts for generation. + batch: A list encoded / tokenized messages (TokensPrompt objects). Returns: The generated text response. @@ -144,6 +169,5 @@ def generate(self, prompts: List[str]) -> List[str]: if self._llm_model is None: raise RuntimeError("Model not initialized") - prompt_tokens = self._apply_chat_template(prompts) - outputs = self._llm_model.generate(prompt_tokens, SamplingParams(**self._generation_kwargs)) + outputs = self._llm_model.generate(batch, SamplingParams(**self._generation_kwargs)) return [output.outputs[0].text for output in outputs] diff --git a/src/eva/language/utils/__init__.py b/src/eva/language/utils/__init__.py index 1e28e9d67..3cded4987 100644 --- a/src/eva/language/utils/__init__.py +++ b/src/eva/language/utils/__init__.py @@ -1,5 +1,6 @@ """Language utilities and helper functions.""" from eva.language.utils.str_to_int_tensor import CastStrToIntTensor +from eva.language.utils.text.messages import format_chat_message -__all__ = ["CastStrToIntTensor"] +__all__ = ["CastStrToIntTensor", "format_chat_message"] diff --git a/src/eva/language/utils/str_to_int_tensor.py b/src/eva/language/utils/str_to_int_tensor.py index dddcfe793..67e2977cd 100644 --- a/src/eva/language/utils/str_to_int_tensor.py +++ b/src/eva/language/utils/str_to_int_tensor.py @@ -16,11 +16,11 @@ class CastStrToIntTensor: Supports single values, lists of strings, or lists of integers. Example: - >>> # Default mapping for yes/no/maybe classification - >>> transform = CastStrToIntTensor() - >>> transform(['yes', 'no', 'maybe']) + >>> # Default mapping for A/B/C classification + >>> transform = CastStrToIntTensor(mapping={"A": 0, "B": 1, "C": 2}) + >>> transform(['B', 'A', 'C']) tensor([1, 0, 2]) - >>> transform('yes') + >>> transform('B') tensor([1]) >>> # Custom mapping @@ -29,20 +29,25 @@ class CastStrToIntTensor: tensor([1, 0]) """ - def __init__(self, mapping: Dict[str, int] | None = None): - """Initialize the transform with a regex-to-integer mapping. + def __init__( + self, mapping: Dict[str, int], standalone_words: bool = True, case_sensitive: bool = True + ) -> None: + r"""Initialize the transform with a regex-to-integer mapping. Args: mapping: Dictionary mapping regex patterns to integers. If None, uses default yes/no/maybe mapping: {'no': 0, 'yes': 1, 'maybe': 2} + standalone_words: If True, patterns are treated as standalone words (e.g., '\bno\b'). + case_sensitive: If True, regex patterns are case-sensitive. """ - if mapping is None: - self.mapping = {r"\bno\b": 0, r"\byes\b": 1, r"\bmaybe\b": 2} - else: - self.mapping = mapping + self.mapping = mapping + + if standalone_words: + self.mapping = {rf"\b{k}\b": v for k, v in mapping.items()} self.compiled_patterns = [ - (re.compile(pattern, re.IGNORECASE), value) for pattern, value in self.mapping.items() + (re.compile(pattern, 0 if case_sensitive else re.IGNORECASE), value) + for pattern, value in self.mapping.items() ] def __call__(self, values: Union[str, List[str], List[int]]) -> torch.Tensor: diff --git a/src/eva/language/utils/text/__init__.py b/src/eva/language/utils/text/__init__.py new file mode 100644 index 000000000..1884d8404 --- /dev/null +++ b/src/eva/language/utils/text/__init__.py @@ -0,0 +1,5 @@ +"""Text utilities for language models.""" + +from eva.language.utils.text.messages import format_chat_message + +__all__ = ["format_chat_message"] diff --git a/src/eva/language/utils/text/messages.py b/src/eva/language/utils/text/messages.py new file mode 100644 index 000000000..89ba753bb --- /dev/null +++ b/src/eva/language/utils/text/messages.py @@ -0,0 +1,67 @@ +"""Message formatting utilities for language models.""" + +import functools +from typing import Any, Dict, List + +from eva.language.data.messages import MessageSeries, SystemMessage + + +def format_chat_message(message: MessageSeries) -> List[Dict[str, Any]]: + """Formats a message series into a format following OpenAI's API specification.""" + return [{"role": item.role, "content": item.content} for item in message] + + +def combine_system_messages(message: MessageSeries, join_char: str = "\n") -> MessageSeries: + """Combine system messages into a single message. + + This is useful when the MessageSeries contains multiple system messages such + as `ModelSystemMessage` and `TaskSystemMessage`. But given that most models / apis + expect a single system message, this function can be used to combines them into one. + + Args: + message: The message series containing one or multiple messages. + join_char: The character to use to join the system messages. Default is newline. + + Returns: + A new message series with system messages combined into one and the + remaining messages unchanged. + """ + system_messages = list(filter(lambda item: item.role == "system", message)) + if len(system_messages) == 0: + return message + + non_system_messages = list(filter(lambda item: item.role != "system", message)) + return [ + SystemMessage(content=merge_message_contents(system_messages, join_char=join_char)) + ] + non_system_messages + + +def merge_message_contents(message: MessageSeries, join_char: str = "\n") -> str: + """Merges the all contents within a message series into a string. + + Args: + message: The message series to combine. + join_char: The character to use to join the message contents. Default is newline. + + Returns: + A string containing the combined message contents. + """ + return join_char.join(item.content for item in message) + + +def insert_system_message( + message: MessageSeries, system_message: SystemMessage | None +) -> MessageSeries: + """Insert a system message at the beginning of the message series.""" + if system_message is None: + return message + return [system_message] + message + + +def batch_insert_system_message( + messages: List[MessageSeries], system_message: SystemMessage | None +) -> List[MessageSeries]: + """Insert a system message at the beginning of each message series in a batch.""" + return list( + map(functools.partial(insert_system_message, system_message=system_message), messages) + ) diff --git a/src/eva/multimodal/__init__.py b/src/eva/multimodal/__init__.py new file mode 100644 index 000000000..ee7b8b43c --- /dev/null +++ b/src/eva/multimodal/__init__.py @@ -0,0 +1,6 @@ +"""Multimodal API.""" + +from eva.multimodal import models +from eva.multimodal.data import datasets + +__all__ = ["models", "datasets"] diff --git a/src/eva/multimodal/data/__init__.py b/src/eva/multimodal/data/__init__.py new file mode 100644 index 000000000..4de0e5b2f --- /dev/null +++ b/src/eva/multimodal/data/__init__.py @@ -0,0 +1,5 @@ +"""Data components for multimodal learning.""" + +from eva.multimodal.data import datasets + +__all__ = ["datasets"] diff --git a/src/eva/multimodal/data/dataloaders/__init__.py b/src/eva/multimodal/data/dataloaders/__init__.py new file mode 100644 index 000000000..c851769c9 --- /dev/null +++ b/src/eva/multimodal/data/dataloaders/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal dataloaders API.""" + +from eva.multimodal.data.dataloaders.collate_fn import text_image_collate + +__all__ = ["text_image_collate"] diff --git a/src/eva/multimodal/data/dataloaders/collate_fn/__init__.py b/src/eva/multimodal/data/dataloaders/collate_fn/__init__.py new file mode 100644 index 000000000..73147ad2e --- /dev/null +++ b/src/eva/multimodal/data/dataloaders/collate_fn/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal collate functions API.""" + +from eva.multimodal.data.dataloaders.collate_fn.text_image import text_image_collate + +__all__ = ["text_image_collate"] diff --git a/src/eva/multimodal/data/dataloaders/collate_fn/text_image.py b/src/eva/multimodal/data/dataloaders/collate_fn/text_image.py new file mode 100644 index 000000000..f49bd4f5b --- /dev/null +++ b/src/eva/multimodal/data/dataloaders/collate_fn/text_image.py @@ -0,0 +1,28 @@ +"""Collate functions for text-image data.""" + +from typing import List + +from torch.utils.data._utils.collate import default_collate + +from eva.multimodal.data.datasets.typings import TextImageSample +from eva.multimodal.models.typings import TextImageBatch + + +def text_image_collate(batch: List[TextImageSample]) -> TextImageBatch: + """Collate function for text-image batches.""" + texts, images, targets, metadata = zip(*batch, strict=False) + + first_sample = batch[0] + metadata = None + if first_sample.metadata is not None: + metadata = { + k: [sample.metadata[k] for sample in batch if sample.metadata] + for k in first_sample.metadata.keys() + } + + return TextImageBatch( + text=list(texts), + image=list(images), + target=default_collate(targets) if targets[0] is not None else None, + metadata=metadata, + ) diff --git a/src/eva/multimodal/data/datasets/__init__.py b/src/eva/multimodal/data/datasets/__init__.py new file mode 100644 index 000000000..cb146c61b --- /dev/null +++ b/src/eva/multimodal/data/datasets/__init__.py @@ -0,0 +1,6 @@ +"""Multimodal datasets API.""" + +from eva.multimodal.data.datasets.multiple_choice.patch_camelyon import PatchCamelyon +from eva.multimodal.data.datasets.text_image import TextImageDataset + +__all__ = ["TextImageDataset", "PatchCamelyon"] diff --git a/src/eva/multimodal/data/datasets/base.py b/src/eva/multimodal/data/datasets/base.py new file mode 100644 index 000000000..ddb5c1426 --- /dev/null +++ b/src/eva/multimodal/data/datasets/base.py @@ -0,0 +1,13 @@ +"""Multimodal Dataset base class.""" + +import abc +from typing import Generic, TypeVar + +from eva.core.data.datasets import base + +DataSample = TypeVar("DataSample") +"""The data sample type.""" + + +class MultimodalDataset(base.MapDataset, abc.ABC, Generic[DataSample]): + """Base dataset class for multimodal tasks.""" diff --git a/src/eva/multimodal/data/datasets/multiple_choice/__init__.py b/src/eva/multimodal/data/datasets/multiple_choice/__init__.py new file mode 100644 index 000000000..a7a4cfc46 --- /dev/null +++ b/src/eva/multimodal/data/datasets/multiple_choice/__init__.py @@ -0,0 +1,5 @@ +"""Multiple choice datasets.""" + +from eva.multimodal.data.datasets.multiple_choice.patch_camelyon import PatchCamelyon + +__all__ = ["PatchCamelyon"] diff --git a/src/eva/multimodal/data/datasets/multiple_choice/patch_camelyon.py b/src/eva/multimodal/data/datasets/multiple_choice/patch_camelyon.py new file mode 100644 index 000000000..5b957537f --- /dev/null +++ b/src/eva/multimodal/data/datasets/multiple_choice/patch_camelyon.py @@ -0,0 +1,80 @@ +"""PatchCamelyon dataset with text prompts for multimodal classification.""" + +from typing import Any, Dict, Literal + +from torchvision import tv_tensors +from typing_extensions import override + +from eva.language.data.messages import MessageSeries, UserMessage +from eva.multimodal.data.datasets.schemas import TransformsSchema +from eva.multimodal.data.datasets.text_image import TextImageDataset +from eva.vision.data import datasets as vision_datasets + + +class PatchCamelyon(TextImageDataset[int], vision_datasets.PatchCamelyon): + """PatchCamelyon image classification using a multiple choice text prompt.""" + + _default_prompt = ( + "You are a pathology expert helping pathologists to analyze images of tissue samples.\n" + "Question: Does this image show metastatic breast tissue?\n" + "Options: A: no, B: yes\n" + "Only answer with a single letter without further explanation. " + "Please always provide an answer, even if you are not sure.\n" + "Answer: " + ) + + def __init__( + self, + root: str, + split: Literal["train", "val", "test"], + download: bool = False, + transforms: TransformsSchema | None = None, + prompt: str | None = None, + max_samples: int | None = None, + ) -> None: + """Initializes the dataset. + + Args: + root: The path to the dataset root. This path should contain + the uncompressed h5 files and the metadata. + split: The dataset split for training, validation, or testing. + download: Whether to download the data for the specified split. + Note that the download will be executed only by additionally + calling the :meth:`prepare_data` method. + transforms: A function/transform which returns a transformed + version of the raw data samples. + prompt: The text prompt to use for classification (multple choice). + max_samples: Maximum number of samples to use. If None, use all samples. + """ + super().__init__(root=root, split=split, download=download, transforms=transforms) + + self.max_samples = max_samples + self.prompt = prompt or self._default_prompt + + if self.max_samples is not None: + self._expected_length = {split: max_samples} + + @property + @override + def class_to_idx(self) -> Dict[str, int]: + return {"A": 0, "B": 1} + + @override + def __len__(self) -> int: + return self.max_samples or self._fetch_dataset_length() + + @override + def load_text(self, index: int) -> MessageSeries: + return [UserMessage(content=self.prompt)] + + @override + def load_image(self, index: int) -> tv_tensors.Image: + return vision_datasets.PatchCamelyon.load_data(self, index) + + @override + def load_target(self, index: int) -> int: + return int(vision_datasets.PatchCamelyon.load_target(self, index).item()) + + @override + def load_metadata(self, index: int) -> Dict[str, Any] | None: + return vision_datasets.PatchCamelyon.load_metadata(self, index) diff --git a/src/eva/multimodal/data/datasets/schemas.py b/src/eva/multimodal/data/datasets/schemas.py new file mode 100644 index 000000000..1934e4618 --- /dev/null +++ b/src/eva/multimodal/data/datasets/schemas.py @@ -0,0 +1,14 @@ +"""Schema definitions for dataset classes.""" + +import dataclasses +from typing import Callable + +from eva.language.data.datasets import schemas as language_schemas + + +@dataclasses.dataclass(frozen=True) +class TransformsSchema(language_schemas.TransformsSchema): + """Schema for dataset transforms.""" + + image: Callable | None = None + """Image transformation""" diff --git a/src/eva/multimodal/data/datasets/text_image.py b/src/eva/multimodal/data/datasets/text_image.py new file mode 100644 index 000000000..57f349f0a --- /dev/null +++ b/src/eva/multimodal/data/datasets/text_image.py @@ -0,0 +1,77 @@ +"""Base classes for text-image datasets.""" + +import abc +from typing import Generic + +from torchvision import tv_tensors +from typing_extensions import override + +from eva.language.data.datasets.text import TextDataset +from eva.multimodal.data.datasets.base import MultimodalDataset +from eva.multimodal.data.datasets.schemas import TransformsSchema +from eva.multimodal.data.datasets.typings import TargetType, TextImageSample + + +class TextImageDataset( + MultimodalDataset[TextImageSample[TargetType]], TextDataset, abc.ABC, Generic[TargetType] +): + """Base dataset class for text-image tasks.""" + + def __init__(self, *args, transforms: TransformsSchema | None = None, **kwargs) -> None: + """Initializes the dataset. + + Args: + *args: Positional arguments for the base class. + transforms: The transforms to apply to the text, image and target when + loading the samples. + **kwargs: Keyword arguments for the base class. + """ + super().__init__(*args, **kwargs) + + self.transforms = transforms + + @abc.abstractmethod + def load_image(self, index: int) -> tv_tensors.Image: + """Returns the image content. + + Args: + index: The index of the data sample. + + Returns: + The image content. + """ + raise NotImplementedError + + @override + def __getitem__(self, index: int) -> TextImageSample[TargetType]: + item = TextImageSample( + text=self.load_text(index), + image=self.load_image(index), + target=self.load_target(index), + metadata=self.load_metadata(index) or {}, + ) + return self._apply_transforms(item) + + @override + def _apply_transforms(self, sample: TextImageSample[TargetType]) -> TextImageSample[TargetType]: + """Applies the dataset transforms to the text, image and target. + + Args: + sample: The sample containing text, image, target and metadata. + + Returns: + The transformed sample. + """ + if self.transforms: + text = self.transforms.text(sample.text) if self.transforms.text else sample.text + image = self.transforms.image(sample.image) if self.transforms.image else sample.image + target = ( + self.transforms.target(sample.target) if self.transforms.target else sample.target + ) + return TextImageSample( + text=text, + image=image, + target=target, + metadata=sample.metadata, + ) + return sample diff --git a/src/eva/multimodal/data/datasets/typings.py b/src/eva/multimodal/data/datasets/typings.py new file mode 100644 index 000000000..4a8fc4796 --- /dev/null +++ b/src/eva/multimodal/data/datasets/typings.py @@ -0,0 +1,27 @@ +"""Typings for multimodal datasets.""" + +from typing import Any, Generic, TypeVar + +from torchvision import tv_tensors +from typing_extensions import NamedTuple + +from eva.language.data.messages import MessageSeries + +TargetType = TypeVar("TargetType") +"""The target data type.""" + + +class TextImageSample(NamedTuple, Generic[TargetType]): + """Text and image sample with target and metadata.""" + + text: MessageSeries + """One or multiple conversation messages.""" + + image: tv_tensors.Image + """Image tensor.""" + + target: TargetType | None + """Target data.""" + + metadata: dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/multimodal/models/__init__.py b/src/eva/multimodal/models/__init__.py new file mode 100644 index 000000000..84d6e30d5 --- /dev/null +++ b/src/eva/multimodal/models/__init__.py @@ -0,0 +1,8 @@ +"""Multimodal models API.""" + +from eva.multimodal.models import networks, wrappers + +__all__ = [ + "networks", + "wrappers", +] diff --git a/src/eva/multimodal/models/modules/__init__.py b/src/eva/multimodal/models/modules/__init__.py new file mode 100644 index 000000000..5f72cc93f --- /dev/null +++ b/src/eva/multimodal/models/modules/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal Networks API.""" + +from eva.multimodal.models.modules.vision_language import VisionLanguageModule + +__all__ = ["VisionLanguageModule"] diff --git a/src/eva/multimodal/models/modules/vision_language.py b/src/eva/multimodal/models/modules/vision_language.py new file mode 100644 index 000000000..5248e5d67 --- /dev/null +++ b/src/eva/multimodal/models/modules/vision_language.py @@ -0,0 +1,55 @@ +"""Model module for vision-language models.""" + +from typing import Any, List + +from lightning.pytorch.utilities.types import STEP_OUTPUT +from torch import nn +from typing_extensions import override + +from eva.core.metrics import structs as metrics_lib +from eva.core.models.modules import module +from eva.core.models.modules.utils import batch_postprocess +from eva.multimodal.models.typings import TextImageBatch + + +class VisionLanguageModule(module.ModelModule): + """Model module for vision-language tasks.""" + + def __init__( + self, + model: nn.Module, + metrics: metrics_lib.MetricsSchema | None = None, + postprocess: batch_postprocess.BatchPostProcess | None = None, + ) -> None: + """Initializes the text inference module. + + Args: + model: Model instance to use for forward pass. + metrics: Metrics schema for evaluation. + postprocess: A helper function to post-process model outputs before evaluation. + """ + super().__init__(metrics=metrics, postprocess=postprocess) + + self.model = model + + @override + def forward(self, batch: TextImageBatch, *args: Any, **kwargs: Any) -> List[str]: + return self.model(batch) + + @override + def validation_step(self, batch: TextImageBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + @override + def test_step(self, batch: TextImageBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + def _batch_step(self, batch: TextImageBatch) -> STEP_OUTPUT: + text, _, targets, metadata = TextImageBatch(*batch) + predictions = self.forward(batch) + return { + "inputs": text, + "predictions": predictions, + "targets": targets, + "metadata": metadata, + } diff --git a/src/eva/multimodal/models/networks/__init__.py b/src/eva/multimodal/models/networks/__init__.py new file mode 100644 index 000000000..59cee4b0f --- /dev/null +++ b/src/eva/multimodal/models/networks/__init__.py @@ -0,0 +1,14 @@ +"""Multimodal networks API.""" + +from eva.multimodal.models.networks.alibaba import Qwen25VL7BInstruct +from eva.multimodal.models.networks.api import Claude35Sonnet20240620, Claude37Sonnet20250219 +from eva.multimodal.models.networks.others import PathoR13b +from eva.multimodal.models.networks.registry import model_registry + +__all__ = [ + "Claude35Sonnet20240620", + "Claude37Sonnet20250219", + "PathoR13b", + "Qwen25VL7BInstruct", + "model_registry", +] diff --git a/src/eva/multimodal/models/networks/alibaba.py b/src/eva/multimodal/models/networks/alibaba.py new file mode 100644 index 000000000..070aaa9d8 --- /dev/null +++ b/src/eva/multimodal/models/networks/alibaba.py @@ -0,0 +1,39 @@ +"""Models from Alibaba.""" + +import torch + +from eva.multimodal.models import wrappers +from eva.multimodal.models.networks.registry import model_registry + + +@model_registry.register("alibaba/qwen2-5-vl-7b-instruct") +class Qwen25VL7BInstruct(wrappers.HuggingFaceModel): + """Qwen2.5-VL 7B Instruct model.""" + + def __init__( + self, + system_prompt: str | None = None, + cache_dir: str | None = None, + attn_implementation: str = "flash_attention_2", + ): + """Initialize the model.""" + super().__init__( + model_name_or_path="Qwen/Qwen2.5-VL-7B-Instruct", + model_class="Qwen2_5_VLForConditionalGeneration", + model_kwargs={ + "torch_dtype": torch.bfloat16, + "trust_remote_code": True, + "cache_dir": cache_dir, + "attn_implementation": attn_implementation, + }, + generation_kwargs={ + "max_new_tokens": 512, + "do_sample": False, + }, + processor_kwargs={ + "padding": True, + "padding_side": "left", + "max_pixels": 451584, # 672*672 + }, + system_prompt=system_prompt, + ) diff --git a/src/eva/multimodal/models/networks/api/__init__.py b/src/eva/multimodal/models/networks/api/__init__.py new file mode 100644 index 000000000..9acac55ca --- /dev/null +++ b/src/eva/multimodal/models/networks/api/__init__.py @@ -0,0 +1,11 @@ +"""Multimodal API networks.""" + +from eva.multimodal.models.networks.api.anthropic import ( + Claude35Sonnet20240620, + Claude37Sonnet20250219, +) + +__all__ = [ + "Claude35Sonnet20240620", + "Claude37Sonnet20250219", +] diff --git a/src/eva/multimodal/models/networks/api/anthropic.py b/src/eva/multimodal/models/networks/api/anthropic.py new file mode 100644 index 000000000..36de34b70 --- /dev/null +++ b/src/eva/multimodal/models/networks/api/anthropic.py @@ -0,0 +1,34 @@ +"""Models from Anthropic.""" + +import os + +from eva.multimodal.models import wrappers +from eva.multimodal.models.networks.registry import model_registry + + +class _Claude(wrappers.LiteLLMModel): + """Base class for Claude models.""" + + def __init__(self, model_name: str, system_prompt: str | None = None): + if not os.getenv("ANTHROPIC_API_KEY"): + raise ValueError("ANTHROPIC_API_KEY env variable must be set.") + + super().__init__(model_name=model_name, system_prompt=system_prompt) + + +@model_registry.register("anthropic/claude-3-5-sonnet-20240620") +class Claude35Sonnet20240620(_Claude): + """Claude 3.5 Sonnet (June 2024) model.""" + + def __init__(self, system_prompt: str | None = None): + """Initialize the model.""" + super().__init__(model_name="claude-3-5-sonnet-20240620", system_prompt=system_prompt) + + +@model_registry.register("anthropic/claude-3-7-sonnet-20250219") +class Claude37Sonnet20250219(_Claude): + """Claude 3.7 Sonnet (February 2025) model.""" + + def __init__(self, system_prompt: str | None = None): + """Initialize the model.""" + super().__init__(model_name="claude-3-7-sonnet-20250219", system_prompt=system_prompt) diff --git a/src/eva/multimodal/models/networks/others.py b/src/eva/multimodal/models/networks/others.py new file mode 100644 index 000000000..05cca2227 --- /dev/null +++ b/src/eva/multimodal/models/networks/others.py @@ -0,0 +1,47 @@ +"""Models from other providers (non-major entities).""" + +import os + +import torch + +from eva.core.utils import requirements +from eva.multimodal.models import wrappers +from eva.multimodal.models.networks.registry import model_registry + + +@model_registry.register("others/wenchuanzhang_patho-r1-3b") +class PathoR13b(wrappers.HuggingFaceModel): + """Patho-R1-3B model by Wenchuan Zhang.""" + + def __init__( + self, + system_prompt: str | None = None, + cache_dir: str | None = None, + attn_implementation: str = "flash_attention_2", + ): + """Initialize the Patho-R1-3B model.""" + requirements.check_dependencies(requirements={"torch": "2.5.1", "torchvision": "0.20.1"}) + + if not os.getenv("HF_TOKEN"): + raise ValueError("HF_TOKEN env variable must be set.") + + super().__init__( + model_name_or_path="WenchuanZhang/Patho-R1-3B", + model_class="Qwen2_5_VLForConditionalGeneration", + model_kwargs={ + "torch_dtype": torch.float16, + "trust_remote_code": True, + "cache_dir": cache_dir, + "attn_implementation": attn_implementation, + }, + generation_kwargs={ + "max_new_tokens": 512, + "do_sample": False, + }, + processor_kwargs={ + "padding": True, + "padding_side": "left", + "max_pixels": 451584, # 672*672 + }, + system_prompt=system_prompt, + ) diff --git a/src/eva/multimodal/models/networks/registry.py b/src/eva/multimodal/models/networks/registry.py new file mode 100644 index 000000000..e370256d6 --- /dev/null +++ b/src/eva/multimodal/models/networks/registry.py @@ -0,0 +1,5 @@ +"""Multimodal Model Registry.""" + +from eva.core.utils.registry import Registry + +model_registry = Registry() diff --git a/src/eva/multimodal/models/typings.py b/src/eva/multimodal/models/typings.py new file mode 100644 index 000000000..82f1bb89c --- /dev/null +++ b/src/eva/multimodal/models/typings.py @@ -0,0 +1,27 @@ +"""Type definitions for multimodal models.""" + +from typing import Any, Dict, Generic, List, TypeVar + +from torchvision import tv_tensors +from typing_extensions import NamedTuple + +from eva.language.data.messages import MessageSeries + +TargetType = TypeVar("TargetType") +"""The target data type.""" + + +class TextImageBatch(NamedTuple, Generic[TargetType]): + """Text and image sample with target and metadata.""" + + text: List[MessageSeries] + """A batch of conversations with one or multiple messages each.""" + + image: List[tv_tensors.Image] + """Image tensor.""" + + target: TargetType | None + """Target data.""" + + metadata: Dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/multimodal/models/wrappers/__init__.py b/src/eva/multimodal/models/wrappers/__init__.py new file mode 100644 index 000000000..5055e02d8 --- /dev/null +++ b/src/eva/multimodal/models/wrappers/__init__.py @@ -0,0 +1,13 @@ +"""Multimodal Wrapper API.""" + +from eva.multimodal.models.wrappers.base import VisionLanguageModel +from eva.multimodal.models.wrappers.from_registry import ModelFromRegistry +from eva.multimodal.models.wrappers.huggingface import HuggingFaceModel +from eva.multimodal.models.wrappers.litellm import LiteLLMModel + +__all__ = [ + "HuggingFaceModel", + "LiteLLMModel", + "ModelFromRegistry", + "VisionLanguageModel", +] diff --git a/src/eva/multimodal/models/wrappers/base.py b/src/eva/multimodal/models/wrappers/base.py new file mode 100644 index 000000000..86bcd3d38 --- /dev/null +++ b/src/eva/multimodal/models/wrappers/base.py @@ -0,0 +1,47 @@ +"""Base class for vision language model wrappers.""" + +import abc +from typing import Any, Callable, List + +from typing_extensions import override + +from eva.core.models.wrappers import base +from eva.language.data.messages import ModelSystemMessage +from eva.multimodal.models.typings import TextImageBatch + + +class VisionLanguageModel(base.BaseModel[TextImageBatch, List[str]]): + """Base class for multimodal models. + + Classes that inherit from this should implement the following methods: + - `load_model`: Loads & instantiates the model. + - `model_forward`: Implements the forward pass of the model. For API models, + this can be an API call. + - `format_inputs`: Preprocesses and converts the input batch into the format + expected by the `model_forward` method. + """ + + def __init__( + self, system_prompt: str | None, output_transforms: Callable | None = None + ) -> None: + """Creates a new model instance. + + Args: + system_prompt: The system prompt to use for the model (optional). + output_transforms: Optional transforms to apply to the output of + the model's forward pass. + """ + super().__init__(transforms=output_transforms) + + self.system_message = ModelSystemMessage(content=system_prompt) if system_prompt else None + + @override + def forward(self, batch: TextImageBatch) -> List[str]: + """Forward pass of the model.""" + inputs = self.format_inputs(batch) + return super().forward(inputs) + + @abc.abstractmethod + def format_inputs(self, batch: TextImageBatch) -> Any: + """Converts the inputs into the format expected by the model.""" + raise NotImplementedError diff --git a/src/eva/multimodal/models/wrappers/from_registry.py b/src/eva/multimodal/models/wrappers/from_registry.py new file mode 100644 index 000000000..1a1e5884e --- /dev/null +++ b/src/eva/multimodal/models/wrappers/from_registry.py @@ -0,0 +1,54 @@ +"""Vision backbone helper class.""" + +from typing import Any, Callable, Dict, List + +from torch import nn +from typing_extensions import override + +from eva.core.models.wrappers import base +from eva.core.utils import factory +from eva.multimodal.models.networks.registry import model_registry +from eva.multimodal.models.typings import TextImageBatch + + +class ModelFromRegistry(base.BaseModel[TextImageBatch, List[str]]): + """Wrapper class for vision backbone models. + + This class can be used by load backbones available in eva's + model registry by name. New backbones can be registered by using + the `@backbone_registry.register(model_name)` decorator. + """ + + def __init__( + self, + model_name: str, + model_kwargs: Dict[str, Any] | None = None, + model_extra_kwargs: Dict[str, Any] | None = None, + transforms: Callable | None = None, + ) -> None: + """Initializes the model. + + Args: + model_name: The name of the model to load. + model_kwargs: The arguments used for instantiating the model. + model_extra_kwargs: Extra arguments used for instantiating the model. + transforms: The transforms to apply to the output tensor + produced by the model. + """ + super().__init__(transforms=transforms) + + self._model_name = model_name + self._model_kwargs = model_kwargs or {} + self._model_extra_kwargs = model_extra_kwargs or {} + + self.model = self.load_model() + + @override + def load_model(self) -> nn.Module: + ModelFromRegistry.__name__ = self._model_name + + return factory.ModuleFactory( + registry=model_registry, + name=self._model_name, + init_args=self._model_kwargs | self._model_extra_kwargs, + ) diff --git a/src/eva/multimodal/models/wrappers/huggingface.py b/src/eva/multimodal/models/wrappers/huggingface.py new file mode 100644 index 000000000..f3c447a70 --- /dev/null +++ b/src/eva/multimodal/models/wrappers/huggingface.py @@ -0,0 +1,180 @@ +"""HuggingFace Vision-Language Model Wrapper.""" + +import functools +from typing import Any, Callable, Dict, List + +import torch +import transformers +from loguru import logger +from torch import nn +from typing_extensions import override + +from eva.language.models.typings import TextBatch +from eva.language.utils.text import messages as language_message_utils +from eva.multimodal.models.typings import TextImageBatch +from eva.multimodal.models.wrappers import base +from eva.multimodal.utils.text import messages as message_utils + + +class HuggingFaceModel(base.VisionLanguageModel): + """Lightweight wrapper for Huggingface VLMs. + + Args: + model_name_or_path: The name of the model to use. + model_class: The class of the model to use. + model_kwargs: Additional model arguments. + processor_kwargs: Additional processor arguments. + generation_kwargs: Additional generation arguments. + """ + + def __init__( + self, + model_name_or_path: str, + model_class: str, + model_kwargs: Dict[str, Any] | None = None, + system_prompt: str | None = None, + processor_kwargs: Dict[str, Any] | None = None, + generation_kwargs: Dict[str, Any] | None = None, + ): + """Initialize the HuggingFace model wrapper. + + Args: + model_name_or_path: The name or path of the model to use. + model_class: The class of the model to use. + model_kwargs: Additional model arguments. + system_prompt: System prompt to use. + processor_kwargs: Additional processor arguments. + generation_kwargs: Additional generation arguments. + """ + super().__init__(system_prompt=system_prompt) + + self.model_name_or_path = model_name_or_path + self.model_kwargs = model_kwargs or {} + self.base_model_class = model_class + self.processor_kwargs = processor_kwargs or {} + self.generation_kwargs = generation_kwargs or {} + + self.processor = self.load_processor() + self.model = self.load_model() + + @override + def format_inputs(self, batch: TextImageBatch | TextBatch) -> Dict[str, torch.Tensor]: + """Formats inputs for HuggingFace models. + + Args: + batch: A batch of text and image inputs. + + Returns: + A dictionary produced by the provided processor following a format like: + { + "input_ids": ..., + "attention_mask": ..., + "pixel_values": ... + } + """ + message_batch, image_batch, _, _ = self._unpack_batch(batch) + with_images = image_batch is not None + + message_batch = language_message_utils.batch_insert_system_message( + message_batch, self.system_message + ) + message_batch = list(map(language_message_utils.combine_system_messages, message_batch)) + + if self.processor.chat_template is not None: # type: ignore + templated_text = [ + self.processor.apply_chat_template( # type: ignore + message, + add_generation_prompt=True, + tokenize=False, + ) + for message in map( + functools.partial( + message_utils.format_huggingface_message, + with_images=with_images, + ), + message_batch, + ) + ] + else: + raise NotImplementedError("Currently only chat models are supported.") + + processor_inputs = { + "text": templated_text, + "return_tensors": "pt", + **self.processor_kwargs, + } + + if with_images: + processor_inputs["image"] = [[image] for image in image_batch] + + return self.processor(**processor_inputs).to(self.model.device) # type: ignore + + @override + def model_forward(self, batch: Dict[str, torch.Tensor]) -> List[str]: + """Generates text output from the model. Is called by the `generate` method. + + Args: + batch: A dictionary containing the input data, which may include: + - "text": List of messages formatted for the model. + - "image": List of image tensors. + + Returns: + A dictionary containing the processed input and the model's output. + """ + output = self.model.generate(**batch, **self.generation_kwargs) # type: ignore + return self._decode_output(output, batch["input_ids"].shape[-1]) + + @override + def load_model(self) -> nn.Module: + """Setting up the model. Used for delayed model initialization. + + Raises: + ValueError: If the model class is not found in transformers or if the model + does not support gradient checkpointing but it is enabled. + """ + logger.info(f"Configuring model: {self.model_name_or_path}") + if hasattr(transformers, self.base_model_class): + model_class = getattr(transformers, self.base_model_class) + else: + raise ValueError(f"Model class {self.base_model_class} not found in transformers") + + model = model_class.from_pretrained(self.model_name_or_path, **self.model_kwargs) + + if not hasattr(model, "generate"): + raise ValueError(f"Model {self.model_name_or_path} does not support generation. ") + + return model + + def load_processor(self) -> Callable: + """Initialize the processor.""" + return transformers.AutoProcessor.from_pretrained( + self.model_name_or_path, + **self.processor_kwargs, + ) + + def _unpack_batch(self, batch: TextImageBatch | TextBatch) -> tuple: + if isinstance(batch, TextImageBatch): + return batch.text, batch.image, batch.target, batch.metadata + return batch.text, None, batch.target, batch.metadata + + def _decode_output(self, output: torch.Tensor, instruction_length: int) -> List[str]: + """Decode the model's batch output to text. + + Args: + output: The raw output from the model. + instruction_length: The length of the instruction in the input. + + Returns: + A list of decoded text responses. + """ + decoded_input = self.processor.batch_decode( # type: ignore + output[:, :instruction_length], skip_special_tokens=True + ) + decoded_output = self.processor.batch_decode( # type: ignore + output[:, instruction_length:], skip_special_tokens=True + ) + + logger.debug(f"Decoded input: {decoded_input}") + logger.debug(f"Decoded output: {decoded_output}") + + return decoded_output diff --git a/src/eva/multimodal/models/wrappers/litellm.py b/src/eva/multimodal/models/wrappers/litellm.py new file mode 100644 index 000000000..e72324eff --- /dev/null +++ b/src/eva/multimodal/models/wrappers/litellm.py @@ -0,0 +1,56 @@ +"""LiteLLM vision-language model wrapper.""" + +import logging +from typing import Any, Dict, List + +from typing_extensions import override + +from eva.language.models import wrappers as language_wrappers +from eva.language.utils.text import messages as language_message_utils +from eva.multimodal.models.typings import TextImageBatch +from eva.multimodal.models.wrappers import base +from eva.multimodal.utils.text import messages as message_utils + + +class LiteLLMModel(base.VisionLanguageModel): + """Wrapper class for LiteLLM vision-language models.""" + + def __init__( + self, + model_name: str, + model_kwargs: Dict[str, Any] | None = None, + system_prompt: str | None = None, + log_level: int | None = logging.INFO, + ): + """Initialize the LiteLLM Wrapper. + + Args: + model_name: The name of the model to use. + model_kwargs: Additional keyword arguments to pass during + generation (e.g., `temperature`, `max_tokens`). + system_prompt: The system prompt to use (optional). + log_level: Optional logging level for LiteLLM. Defaults to WARNING. + """ + super().__init__(system_prompt=system_prompt) + + self.language_model = language_wrappers.LiteLLMModel( + model_name=model_name, + model_kwargs=model_kwargs, + system_prompt=system_prompt, + log_level=log_level, + ) + + @override + def format_inputs(self, batch: TextImageBatch) -> List[List[Dict[str, Any]]]: + message_batch, image_batch, _, _ = TextImageBatch(*batch) + + message_batch = language_message_utils.batch_insert_system_message( + message_batch, self.system_message + ) + message_batch = list(map(language_message_utils.combine_system_messages, message_batch)) + + return list(map(message_utils.format_litellm_message, message_batch, image_batch)) + + @override + def model_forward(self, batch: List[List[Dict[str, Any]]]) -> List[str]: + return self.language_model.model_forward(batch) diff --git a/src/eva/multimodal/utils/__init__.py b/src/eva/multimodal/utils/__init__.py new file mode 100644 index 000000000..cf0f30d4a --- /dev/null +++ b/src/eva/multimodal/utils/__init__.py @@ -0,0 +1 @@ +"""Multimodal utilities API.""" diff --git a/src/eva/multimodal/utils/image/__init__.py b/src/eva/multimodal/utils/image/__init__.py new file mode 100644 index 000000000..52fdb6d3d --- /dev/null +++ b/src/eva/multimodal/utils/image/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal image utilities API.""" + +from eva.multimodal.utils.image.encode import encode_image + +__all__ = ["encode_image"] diff --git a/src/eva/multimodal/utils/image/encode.py b/src/eva/multimodal/utils/image/encode.py new file mode 100644 index 000000000..4e71b5568 --- /dev/null +++ b/src/eva/multimodal/utils/image/encode.py @@ -0,0 +1,28 @@ +"""Image encoding utilities.""" + +import base64 +import io +from typing import Literal + +from torchvision import tv_tensors +from torchvision.transforms.v2 import functional as F + + +def encode_image(image: tv_tensors.Image, encoding: Literal["base64"]) -> str: + """Encodes an image tensor into a string format. + + Args: + image: The image tensor to encode. + encoding: The encoding format to use. Currently only supports "base64". + + Returns: + An encoded string representation of the image. + """ + match encoding: + case "base64": + image_bytes = io.BytesIO() + F.to_pil_image(image).save(image_bytes, format="PNG", optimize=True) + image_bytes.seek(0) + return base64.b64encode(image_bytes.getvalue()).decode("utf-8") + case _: + raise ValueError(f"Unsupported encoding type: {encoding}. Supported: 'base64'") diff --git a/src/eva/multimodal/utils/text/__init__.py b/src/eva/multimodal/utils/text/__init__.py new file mode 100644 index 000000000..f4e5a98ae --- /dev/null +++ b/src/eva/multimodal/utils/text/__init__.py @@ -0,0 +1 @@ +"""Multimodal text utilities API.""" diff --git a/src/eva/multimodal/utils/text/messages.py b/src/eva/multimodal/utils/text/messages.py new file mode 100644 index 000000000..573e9cad1 --- /dev/null +++ b/src/eva/multimodal/utils/text/messages.py @@ -0,0 +1,79 @@ +"""Message formatting utilities for multimodal models.""" + +from typing import Any, Dict, List + +from torchvision import tv_tensors + +from eva.language import utils as language_utils +from eva.language.data.messages import MessageSeries +from eva.multimodal.utils import image as image_utils + + +def format_huggingface_message( + message: MessageSeries, with_images: bool = False +) -> List[Dict[str, Any]]: + """Formats a message series into a format suitable for Huggingface models.""" + if not with_images: + return language_utils.format_chat_message(message) + + formatted_message = [] + for item in message: + if item.role == "system": + formatted_message += language_utils.format_chat_message([item]) + else: + formatted_message.append( + { + "role": item.role, + "content": [ + { + "type": "text", + "text": str(item.content), + }, + {"type": "image"}, + ], + } + ) + return formatted_message + + +def format_litellm_message( + message: MessageSeries, image: tv_tensors.Image | None +) -> List[Dict[str, Any]]: + """Format a message series for LiteLLM API. + + Args: + message: The message series to format. + image: Optional image to include in the message. + + Returns: + A list of formatted message dictionaries. + """ + if image is None: + return language_utils.format_chat_message(message) + + formatted_message = [] + for item in message: + if item.role == "system": + formatted_message += language_utils.format_chat_message([item]) + else: + formatted_message.append( + { + "role": item.role, + "content": [ + { + "type": "text", + "text": str(item.content), + }, + { + "type": "image_url", + "image_url": { + "url": ( + f"data:image/png;base64," + f"{image_utils.encode_image(image, encoding='base64')}" + ) + }, + }, + ], + } + ) + return formatted_message diff --git a/src/eva/vision/data/datasets/classification/patch_camelyon.py b/src/eva/vision/data/datasets/classification/patch_camelyon.py index 8d66bce48..92c5fd5d1 100644 --- a/src/eva/vision/data/datasets/classification/patch_camelyon.py +++ b/src/eva/vision/data/datasets/classification/patch_camelyon.py @@ -61,6 +61,13 @@ class PatchCamelyon(vision.VisionDataset[tv_tensors.Image, torch.Tensor]): ] """Test resources.""" + _expected_length = { + "train": 262144, + "val": 32768, + "test": 32768, + } + """Expected dataset length for each split.""" + _license: str = ( "Creative Commons Zero v1.0 Universal (https://choosealicense.com/licenses/cc0-1.0/)" ) @@ -113,14 +120,9 @@ def prepare_data(self) -> None: @override def validate(self) -> None: - expected_length = { - "train": 262144, - "val": 32768, - "test": 32768, - } _validators.check_dataset_integrity( self, - length=expected_length.get(self._split, 0), + length=self._expected_length.get(self._split, 0), n_classes=2, first_and_last_labels=("no_tumor", "tumor"), ) diff --git a/src/eva/vision/data/transforms/__init__.py b/src/eva/vision/data/transforms/__init__.py index 04e276e1a..c2a6f634e 100644 --- a/src/eva/vision/data/transforms/__init__.py +++ b/src/eva/vision/data/transforms/__init__.py @@ -13,10 +13,11 @@ RandShiftIntensity, ScaleIntensityRange, ) -from eva.vision.data.transforms.spatial import RandFlip, RandRotate90, Spacing +from eva.vision.data.transforms.spatial import RandFlip, RandRotate90, Resize, Spacing from eva.vision.data.transforms.utility import EnsureChannelFirst __all__ = [ + "Resize", "ResizeAndCrop", "Squeeze", "CropForeground", diff --git a/src/eva/vision/data/transforms/spatial/__init__.py b/src/eva/vision/data/transforms/spatial/__init__.py index ed3bf4691..51db5ecb4 100644 --- a/src/eva/vision/data/transforms/spatial/__init__.py +++ b/src/eva/vision/data/transforms/spatial/__init__.py @@ -1,7 +1,8 @@ """Transforms for spatial operations.""" from eva.vision.data.transforms.spatial.flip import RandFlip +from eva.vision.data.transforms.spatial.resize import Resize from eva.vision.data.transforms.spatial.rotate import RandRotate90 from eva.vision.data.transforms.spatial.spacing import Spacing -__all__ = ["Spacing", "RandFlip", "RandRotate90"] +__all__ = ["Spacing", "RandFlip", "RandRotate90", "Resize"] diff --git a/src/eva/vision/data/transforms/spatial/functional/__init__.py b/src/eva/vision/data/transforms/spatial/functional/__init__.py new file mode 100644 index 000000000..e2a53aed3 --- /dev/null +++ b/src/eva/vision/data/transforms/spatial/functional/__init__.py @@ -0,0 +1,5 @@ +"""Functional API for spatial transforms.""" + +from eva.vision.data.transforms.spatial.functional.resize import resize_to_max_bytes + +__all__ = ["resize_to_max_bytes"] diff --git a/src/eva/vision/data/transforms/spatial/functional/resize.py b/src/eva/vision/data/transforms/spatial/functional/resize.py new file mode 100644 index 000000000..2eee119cc --- /dev/null +++ b/src/eva/vision/data/transforms/spatial/functional/resize.py @@ -0,0 +1,26 @@ +"""Functional resizing utilities.""" + +import io +from typing import Tuple + +from PIL import Image +from torchvision import tv_tensors +from torchvision.transforms.v2 import functional as F + + +def resize_to_max_bytes(image: tv_tensors.Image, max_bytes: int) -> tv_tensors.Image: + """Resize the image to fit within the specified byte size.""" + image_pil = F.to_pil_image(image) + image_bytes = io.BytesIO() + image_pil.save(image_bytes, format="PNG", optimize=True) + + while image_bytes.tell() > max_bytes: + size: Tuple[int, int] = image_pil.size # type: ignore + w, h = size + scale = (max_bytes / image_bytes.tell()) ** 0.5 + new_size = (max(1, int(h * scale)), max(1, int(w * scale))) + image_pil = image_pil.resize(new_size, Image.Resampling.LANCZOS) + image_bytes = io.BytesIO() + image_pil.save(image_bytes, format="PNG", optimize=True) + + return tv_tensors.Image(F.pil_to_tensor(image_pil)) diff --git a/src/eva/vision/data/transforms/spatial/resize.py b/src/eva/vision/data/transforms/spatial/resize.py new file mode 100644 index 000000000..4678278eb --- /dev/null +++ b/src/eva/vision/data/transforms/spatial/resize.py @@ -0,0 +1,62 @@ +"""Image resize transforms.""" + +import functools +from typing import Any, Dict + +from torchvision import tv_tensors +from torchvision.transforms import v2 +from typing_extensions import override + +from eva.vision.data.transforms.spatial import functional + + +class Resize(v2.Transform): + """Resize transform for images with spatial or byte-based constraints. + + This transform provides two mutually exclusive modes of resizing: + 1. Spatial resizing: Resize to a specific (height, width) dimension + 2. Byte-based resizing: Resize to fit within a maximum byte size + + The latter is particularly useful for API models (e.g. Claude 3.7) that + have strict byte size limits for image inputs. + """ + + def __init__(self, size: tuple[int, int] | None = None, max_bytes: int | None = None) -> None: + """Initializes the transform. + + Args: + size: Target size as (height, width) tuple for spatial resizing. + If provided, max_bytes must be None. + max_bytes: Maximum allowed byte size for the image. + If provided, size must be None. Must be a positive integer. + + Raises: + ValueError: If both size and max_bytes are provided, or if max_bytes + is not a positive integer. + """ + if size is not None and max_bytes is not None: + raise ValueError("Cannot provide both 'size' and 'max_bytes' parameters.") + if max_bytes is not None and max_bytes <= 0: + raise ValueError("'max_bytes' must be a positive integer.") + + super().__init__() + + self.size = size + self.max_bytes = max_bytes + self.resize_fn = None + + if size is not None: + self.resize_fn = v2.Resize(size=size) + elif max_bytes is not None: + self.resize_fn = functools.partial(functional.resize_to_max_bytes, max_bytes=max_bytes) + + @functools.singledispatchmethod + @override + def _transform(self, inpt: Any, params: Dict[str, Any]) -> Any: + return inpt + + @_transform.register(tv_tensors.Image) + @_transform.register(tv_tensors.Mask) + def _(self, inpt: Any, params: Dict[str, Any]) -> Any: + inpt_resized = self.resize_fn(inpt) if self.resize_fn is not None else inpt + return tv_tensors.wrap(inpt_resized, like=inpt) diff --git a/src/eva/vision/models/wrappers/from_registry.py b/src/eva/vision/models/wrappers/from_registry.py index b7529198f..00ad4f142 100644 --- a/src/eva/vision/models/wrappers/from_registry.py +++ b/src/eva/vision/models/wrappers/from_registry.py @@ -3,6 +3,7 @@ from typing import Any, Callable, Dict import torch +from torch import nn from typing_extensions import override from eva.core.models.wrappers import base @@ -40,14 +41,14 @@ def __init__( self._model_kwargs = model_kwargs or {} self._model_extra_kwargs = model_extra_kwargs or {} - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: - self._model = factory.ModuleFactory( + def load_model(self) -> nn.Module: + ModelFromRegistry.__name__ = self._model_name + + return factory.ModuleFactory( registry=backbone_registry, name=self._model_name, init_args=self._model_kwargs | self._model_extra_kwargs, ) - - ModelFromRegistry.__name__ = self._model_name diff --git a/src/eva/vision/models/wrappers/from_timm.py b/src/eva/vision/models/wrappers/from_timm.py index 8bcb773bc..236d9b5c6 100644 --- a/src/eva/vision/models/wrappers/from_timm.py +++ b/src/eva/vision/models/wrappers/from_timm.py @@ -5,6 +5,7 @@ import timm import torch +from torch import nn from typing_extensions import override from eva.core.models.wrappers import base @@ -46,12 +47,14 @@ def __init__( self._out_indices = out_indices self._model_kwargs = model_kwargs or {} - self.load_model() + self.model = self.load_model() @override - def load_model(self) -> None: + def load_model(self) -> nn.Module: """Builds and loads the timm model as feature extractor.""" - self._model = timm.create_model( + TimmModel.__name__ = self._model_name + + return timm.create_model( model_name=self._model_name, pretrained=True if self._checkpoint_path else self._pretrained, pretrained_cfg=self._pretrained_cfg, @@ -59,7 +62,6 @@ def load_model(self) -> None: features_only=self._out_indices is not None, **self._model_kwargs, ) - TimmModel.__name__ = self._model_name @property def _pretrained_cfg(self) -> Dict[str, Any]: diff --git a/tests/__init__.py b/tests/__init__.py index fe3abebdc..62bbab35b 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -1 +1 @@ -"""EVA tests.""" +"""eva tests.""" diff --git a/tests/eva/core/test_cli.py b/tests/eva/core/test_cli.py index ca9a2c64a..cbc53d86d 100644 --- a/tests/eva/core/test_cli.py +++ b/tests/eva/core/test_cli.py @@ -1,4 +1,4 @@ -"""Tests regarding the EVA `fit` CLI command on core datasets.""" +"""Tests regarding the eva `fit` CLI command on core datasets.""" import os diff --git a/tests/eva/language/__init__.py b/tests/eva/language/__init__.py index 7cdb9f6f2..ca010d298 100644 --- a/tests/eva/language/__init__.py +++ b/tests/eva/language/__init__.py @@ -1 +1 @@ -"""EVA language tests.""" +"""eva language tests.""" diff --git a/tests/eva/language/data/datasets/classification/test_pubmedqa.py b/tests/eva/language/data/datasets/classification/test_pubmedqa.py index e07f3becc..4e5ef84af 100644 --- a/tests/eva/language/data/datasets/classification/test_pubmedqa.py +++ b/tests/eva/language/data/datasets/classification/test_pubmedqa.py @@ -8,6 +8,7 @@ from datasets import Dataset from eva.language.data import datasets +from eva.language.data.messages import Message @pytest.mark.parametrize( @@ -35,10 +36,12 @@ def test_sample(pubmedqa_dataset: datasets.PubMedQA, index: int) -> None: assert isinstance(sample, tuple) assert len(sample) == 3 - text, target, metadata = sample - assert isinstance(text, str) - assert text.startswith("Question: ") - assert "Context: " in text + messages, target, metadata = sample + assert isinstance(messages, list) + assert all(isinstance(item, Message) for item in messages) + + assert messages[0].content.startswith("Question: ") + assert "Context: " in messages[0].content assert isinstance(target, torch.Tensor) assert target in [0, 1, 2] diff --git a/tests/eva/language/models/modules/test_language.py b/tests/eva/language/models/modules/test_language.py new file mode 100644 index 000000000..f7046ee23 --- /dev/null +++ b/tests/eva/language/models/modules/test_language.py @@ -0,0 +1,62 @@ +"""Tests the language model module.""" + +import pytest +from torch import nn + +from eva.language.models import LanguageModule +from eva.language.models.typings import TextBatch + + +def test_forward(language_module): + """Test the forward method of the LanguageModule class.""" + input_text = ["Hello world"] + expected = ["Dummy response Nr. 0"] + result = language_module.forward(input_text) + assert result == expected + + +def test_validation_step(language_module, model): + """Test the validation_step method of the LanguageModule class.""" + data = ["What is the capital of France?"] + targets = ["Paris"] + metadata = [{"id": 1}] + batch = (data, targets, metadata) + + # The module creates messages list: [str(d) + "\n" + prompt for d in data] + expected_messages = [f"Dummy response Nr. {i}" for i in range(len(batch))] + expected_predictions = model(expected_messages) + + output = language_module.validation_step(batch) + + assert "predictions" in output + assert "targets" in output + assert "metadata" in output + assert output["predictions"] == expected_predictions + assert output["targets"] == targets + assert output["metadata"] == metadata + + +def test_init_attributes(model): + """Test the attributes of the LanguageModule class.""" + module_instance = LanguageModule(model=model) + assert module_instance.model is model + + +class DummyModel(nn.Module): + """A simple text model for testing purposes.""" + + def forward(self, batch: TextBatch) -> list[str]: + """Generate some text based on the input prompt.""" + return [f"Dummy response Nr. {i}" for i in range(len(batch))] + + +@pytest.fixture +def model(): + """Return a dummy model instance.""" + return DummyModel() + + +@pytest.fixture +def language_module(model): + """Return a LanguageModule instance.""" + return LanguageModule(model=model) diff --git a/tests/eva/language/models/modules/test_text.py b/tests/eva/language/models/modules/test_text.py deleted file mode 100644 index 81b2d96ff..000000000 --- a/tests/eva/language/models/modules/test_text.py +++ /dev/null @@ -1,69 +0,0 @@ -"""Tests the TextModule module.""" - -import pytest -from torch import nn - -from eva.language.models import TextModule - - -def test_forward(text_module, text_model): - """Test the forward method of the TextModule class.""" - input_text = "Hello world" - expected = text_model(input_text) - result = text_module.forward(input_text) - assert result == expected - - -def test_validation_step(text_module, text_model): - """Test the validation_step method of the TextModule class.""" - data = ["What is the capital of France?"] - targets = ["Paris"] - metadata = [{"id": 1}] - batch = (data, targets, metadata) - - # The module creates messages list: [str(d) + "\n" + prompt for d in data] - expected_messages = [str(data[0]) + "\n" + text_module.prompt] - expected_predictions = text_model(expected_messages) - - output = text_module.validation_step(batch) - - assert "predictions" in output - assert "targets" in output - assert "metadata" in output - assert output["predictions"] == expected_predictions - assert output["targets"] == targets - assert output["metadata"] == metadata - - -def test_init_attributes(text_model): - """Test the attributes of the TextModule class.""" - prompt = "Initialization Prompt: " - module_instance = TextModule(model=text_model, prompt=prompt) - assert module_instance.model is text_model - assert module_instance.prompt == prompt - - -class TextModel(nn.Module): - """A simple text model for testing purposes.""" - - def forward(self, prompts): - """Generate some text based on the input prompt.""" - if isinstance(prompts, str): - return [f"Generated: {prompts}"] - elif isinstance(prompts, list): - return [f"Generated: {prompt}" for prompt in prompts] - else: - return [f"Generated: {str(prompts)}"] - - -@pytest.fixture -def text_model(): - """Return a TextModel instance.""" - return TextModel() - - -@pytest.fixture -def text_module(text_model): - """Return a TextModule instance.""" - prompt = "Test Prompt: " - return TextModule(model=text_model, prompt=prompt) diff --git a/tests/eva/language/models/wrappers/test_huggingface.py b/tests/eva/language/models/wrappers/test_huggingface.py index 454de9c84..1e3bff4ec 100644 --- a/tests/eva/language/models/wrappers/test_huggingface.py +++ b/tests/eva/language/models/wrappers/test_huggingface.py @@ -4,7 +4,9 @@ import pytest -from eva.language.models import HuggingFaceTextModel +from eva.language.data.messages import UserMessage +from eva.language.models import HuggingFaceModel +from eva.language.models.typings import TextBatch @pytest.mark.parametrize( @@ -46,14 +48,15 @@ def test_real_small_hf_model_generation( ] with patch("eva.language.models.wrappers.huggingface.pipeline", return_value=mock_pipeline): - model = HuggingFaceTextModel( + model = HuggingFaceModel( model_name_or_path=model_name_or_path, task="text-generation", generation_kwargs=generate_kwargs, ) - output1 = model([prompt])[0] - output2 = model([prompt])[0] + batch = TextBatch(text=[[UserMessage(content=prompt)]], target=None, metadata={}) + output1 = model(batch)[0] + output2 = model(batch)[0] assert isinstance(output1, str) and output1, "First output should be a non-empty string." assert isinstance(output2, str) and output2, "Second output should be a non-empty string." diff --git a/tests/eva/language/models/wrappers/test_litellm.py b/tests/eva/language/models/wrappers/test_litellm.py index 68ca64cb9..89d56d74c 100644 --- a/tests/eva/language/models/wrappers/test_litellm.py +++ b/tests/eva/language/models/wrappers/test_litellm.py @@ -2,15 +2,17 @@ import pytest -from eva.language.models import LiteLLMTextModel +from eva.language.data.messages import UserMessage +from eva.language.models import LiteLLMModel +from eva.language.models.typings import TextBatch -DUMMY_RESPONSE = {"choices": [{"message": {"content": "Test response"}}]} +DUMMY_RESPONSE = {"choices": [{"message": {"content": "Test response", "role": "assistant"}}]} def test_generate(model_instance): """Test that the generate method returns the expected dummy response.""" - prompts = ["Hello, world!"] - result = model_instance(prompts) + batch = TextBatch(text=[[UserMessage(content="Hello, world!")]], target=None, metadata={}) + result = model_instance(batch) assert result == ["Test response"] @@ -31,9 +33,9 @@ def _fake_batch_completion(**_kwargs): @pytest.fixture def model_instance(fake_completion): # noqa: ARG001 - """Fixture to instantiate the LiteLLMTextModel with a valid model name. + """Fixture to instantiate the LiteLLMModel with a valid model name. Using a valid model name (like 'openai/gpt-3.5-turbo') helps pass provider lookup. fake_completion dependency ensures mocking is set up before model creation. """ - return LiteLLMTextModel("openai/gpt-3.5-turbo", model_kwargs={"temperature": 0.7}) + return LiteLLMModel("openai/gpt-3.5-turbo", model_kwargs={"temperature": 0.7}) diff --git a/tests/eva/language/models/wrappers/test_vllm.py b/tests/eva/language/models/wrappers/test_vllm.py index f209e4fda..a4ba98d61 100644 --- a/tests/eva/language/models/wrappers/test_vllm.py +++ b/tests/eva/language/models/wrappers/test_vllm.py @@ -3,7 +3,7 @@ import pytest try: - from eva.language.models.wrappers.vllm import VLLMTextModel + from eva.language.models.wrappers.vllm import VllmModel except ImportError: pytest.skip("vLLM not available", allow_module_level=True) @@ -81,8 +81,8 @@ def mock_vllm_imports(monkeypatch): def test_initialization(mock_vllm_imports): - """Tests VLLMTextModel initialization.""" - model = VLLMTextModel( + """Tests VllmModel initialization.""" + model = VllmModel( model_name_or_path="test/model", model_kwargs={"max_model_len": 1024}, generation_kwargs={"max_tokens": 100}, @@ -95,7 +95,7 @@ def test_initialization(mock_vllm_imports): def test_lazy_loading(mock_vllm_imports): """Tests lazy model loading.""" - model = VLLMTextModel("test/model") + model = VllmModel("test/model") assert model._llm_model is None model.load_model() @@ -105,7 +105,7 @@ def test_lazy_loading(mock_vllm_imports): def test_generate(mock_vllm_imports): """Tests text generation.""" - model = VLLMTextModel("test/model") + model = VllmModel("test/model") prompts = ["Hello", "How are you?"] results = model(prompts) @@ -116,7 +116,7 @@ def test_generate(mock_vllm_imports): def test_chat_template_application(mock_vllm_imports): """Tests chat template application.""" - model = VLLMTextModel("test/model") + model = VllmModel("test/model") prompts = ["Hello world"] token_prompts = model._apply_chat_template(prompts) @@ -137,7 +137,7 @@ def apply_chat_template( def mock_get_tokenizer(): return MockTokenizerDoubleBOS() - model = VLLMTextModel("test/model") + model = VllmModel("test/model") model.load_model() monkeypatch.setattr(model._llm_model, "get_tokenizer", mock_get_tokenizer) model._llm_tokenizer = mock_get_tokenizer() @@ -165,7 +165,7 @@ def __init__(self): """Initialize tokenizer without chat template.""" pass - model = VLLMTextModel("test/model") + model = VllmModel("test/model") model.load_model() monkeypatch.setattr(model._llm_model, "get_tokenizer", lambda: MockTokenizerNoTemplate()) model._llm_tokenizer = MockTokenizerNoTemplate() diff --git a/tests/eva/language/test_language_cli.py b/tests/eva/language/test_language_cli.py index b45199330..d8ceb68f5 100644 --- a/tests/eva/language/test_language_cli.py +++ b/tests/eva/language/test_language_cli.py @@ -14,7 +14,7 @@ @pytest.mark.parametrize( "configuration_file", [ - "configs/language/pubmedqa.yaml", + "configs/language/pathology/online/multiple_choice/pubmedqa.yaml", ], ) def test_configuration_initialization(configuration_file: str, lib_path: str) -> None: @@ -32,7 +32,7 @@ def test_configuration_initialization(configuration_file: str, lib_path: str) -> @pytest.mark.parametrize( "configuration_file", [ - "configs/language/pubmedqa.yaml", + "configs/language/pathology/online/multiple_choice/pubmedqa.yaml", ], ) def test_validate_from_configuration(configuration_file: str, lib_path: str) -> None: @@ -52,7 +52,7 @@ def mock_dependencies(): """Mocks external dependencies to avoid API calls and downloads.""" def _fake_completion(_model, _messages, **_kwargs): - return {"choices": [{"message": {"content": "yes"}}]} + return {"choices": [{"message": {"content": "yes", "role": "assistant"}}]} def _fake_prepare_data(self): # Create a minimal fake dataset matching PubMedQA format @@ -71,5 +71,6 @@ def _fake_prepare_data(self): lambda **_kwargs: [_fake_completion(None, None)], ), mock.patch.dict(os.environ, {"OPENAI_API_KEY": "dummy-key"}), + mock.patch.dict(os.environ, {"ANTHROPIC_API_KEY": "dummy-key"}), ): yield diff --git a/tests/eva/language/utils/test_str_to_int_tensor.py b/tests/eva/language/utils/test_str_to_int_tensor.py index da99a4f2c..7fe2987a2 100644 --- a/tests/eva/language/utils/test_str_to_int_tensor.py +++ b/tests/eva/language/utils/test_str_to_int_tensor.py @@ -5,6 +5,8 @@ from eva.language.utils.str_to_int_tensor import CastStrToIntTensor +DEFAULT_MAPPING = {"no": 0, "yes": 1, "maybe": 2} + @pytest.mark.parametrize( "input_values, expected", @@ -23,7 +25,7 @@ ) def test_cast_str_to_int_tensor_valid(input_values, expected): """Test CastStrToIntTensor with valid inputs.""" - result = CastStrToIntTensor()(input_values) + result = CastStrToIntTensor(mapping=DEFAULT_MAPPING)(input_values) assert torch.equal(result, expected) @@ -34,36 +36,42 @@ def test_cast_str_to_int_tensor_valid(input_values, expected): def test_cast_str_to_int_tensor_invalid(invalid_input): """Test CastStrToIntTensor with invalid inputs.""" with pytest.raises(ValueError): - CastStrToIntTensor()(invalid_input) + CastStrToIntTensor(mapping=DEFAULT_MAPPING)(invalid_input) @pytest.mark.parametrize( - "custom_mapping, input_values, expected", + "custom_mapping, input_values, case_sensitive, expected", [ ( {r"positive|good": 1, r"negative|bad": 0}, ["positive", "bad"], + True, torch.tensor([1, 0], dtype=torch.int), ), ( {r"positive|good": 1, r"negative|bad": 0}, ["good", "negative"], + True, torch.tensor([1, 0], dtype=torch.int), ), ( {r"positive|good": 1, r"negative|bad": 0}, ["POSITIVE", "BAD"], + False, torch.tensor([1, 0], dtype=torch.int), ), ( {r"\bhappy\b": 1, r"\bsad\b": 0}, ["I am happy", "feeling sad"], + True, torch.tensor([1, 0], dtype=torch.int), ), ], ) -def test_cast_str_to_int_tensor_custom_mapping(custom_mapping, input_values, expected): +def test_cast_str_to_int_tensor_custom_mapping( + custom_mapping: dict, input_values: list, case_sensitive: bool, expected: torch.Tensor +): """Test CastStrToIntTensor with custom mapping.""" - transform = CastStrToIntTensor(custom_mapping) + transform = CastStrToIntTensor(mapping=custom_mapping, case_sensitive=case_sensitive) result = transform(input_values) assert torch.equal(result, expected) diff --git a/tests/eva/multimodal/__init__.py b/tests/eva/multimodal/__init__.py new file mode 100644 index 000000000..0c3ee61ad --- /dev/null +++ b/tests/eva/multimodal/__init__.py @@ -0,0 +1 @@ +"""eva multimodal tests.""" diff --git a/tests/eva/multimodal/test_multimodal_cli.py b/tests/eva/multimodal/test_multimodal_cli.py new file mode 100644 index 000000000..a63a80592 --- /dev/null +++ b/tests/eva/multimodal/test_multimodal_cli.py @@ -0,0 +1,72 @@ +"""Tests regarding eva's CLI commands on multimodal datasets.""" + +import os +from unittest import mock +from unittest.mock import patch + +import pytest + +from eva.multimodal.data import datasets +from tests.eva import _cli + +BATCH_SIZE = 2 + + +@pytest.mark.parametrize( + "configuration_file", + [ + "configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml", + ], +) +def test_configuration_initialization(configuration_file: str, lib_path: str) -> None: + """Tests that a given configuration file can be initialized.""" + _cli.run_cli_from_main( + cli_args=[ + "validate", + "--config", + os.path.join(lib_path, configuration_file), + "--print_config", + ] + ) + + +@pytest.mark.parametrize( + "configuration_file", + [ + "configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml", + ], +) +def test_validate_from_configuration(configuration_file: str, lib_path: str) -> None: + """Tests CLI `validate` command with a given configuration file.""" + with mock.patch.dict(os.environ, {"N_RUNS": "1", "BATCH_SIZE": f"{BATCH_SIZE}"}): + _cli.run_cli_from_main( + cli_args=[ + "validate", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + + +@pytest.fixture(autouse=True) +def skip_dataset_validation() -> None: + """Mocks the validation step of the datasets.""" + datasets.PatchCamelyon.validate = mock.MagicMock(return_value=None) + + +@pytest.fixture(autouse=True) +def mock_dependencies(): + """Mocks external dependencies to avoid API calls and downloads.""" + + def _fake_completion(): + return {"choices": [{"message": {"content": "A", "role": "assistant"}}]} + + with ( + patch( + "eva.language.models.wrappers.litellm.batch_completion", + lambda **_kwargs: [_fake_completion()] * BATCH_SIZE, + ), + mock.patch.dict(os.environ, {"OPENAI_API_KEY": "dummy-key"}), + mock.patch.dict(os.environ, {"ANTHROPIC_API_KEY": "dummy-key"}), + ): + yield diff --git a/tests/eva/vision/__init__.py b/tests/eva/vision/__init__.py index 6a9dd35ed..367c6e8fa 100644 --- a/tests/eva/vision/__init__.py +++ b/tests/eva/vision/__init__.py @@ -1 +1 @@ -"""EVA vision tests.""" +"""eva vision tests.""" From bd1be76dc8876163949dd208af59538036832e8a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20K=C3=A4nzig?= <36882833+nkaenzig@users.noreply.github.com> Date: Fri, 29 Aug 2025 17:01:11 +0200 Subject: [PATCH 2/3] Add unit tests for `eva.multimodal` (#871) --- .../data/dataloaders/collate_fn/test_text.py | 182 +++++++++++++ tests/eva/language/data/datasets/test_text.py | 80 ++++++ .../models/wrappers/test_from_registry.py | 28 ++ tests/eva/language/utils/test_messages.py | 241 ++++++++++++++++++ tests/eva/multimodal/data/__init__.py | 1 + .../multimodal/data/dataloaders/__init__.py | 1 + .../data/dataloaders/collate_fn/__init__.py | 1 + .../dataloaders/collate_fn/test_text_image.py | 88 +++++++ .../eva/multimodal/data/datasets/__init__.py | 1 + .../data/datasets/multiple_choice/__init__.py | 1 + .../multiple_choice/test_patch_camelyon.py | 104 ++++++++ .../eva/multimodal/data/datasets/test_base.py | 31 +++ .../data/datasets/test_dataset_typings.py | 50 ++++ .../multimodal/data/datasets/test_schemas.py | 48 ++++ .../data/datasets/test_text_image.py | 90 +++++++ tests/eva/multimodal/models/__init__.py | 1 + .../models/modules/test_vision_language.py | 113 ++++++++ tests/eva/multimodal/models/test_typings.py | 36 +++ .../multimodal/models/wrappers/__init__.py | 1 + .../multimodal/models/wrappers/test_base.py | 46 ++++ .../models/wrappers/test_from_registry.py | 28 ++ .../models/wrappers/test_huggingface.py | 124 +++++++++ .../models/wrappers/test_litellm.py | 89 +++++++ tests/eva/multimodal/utils/__init__.py | 1 + tests/eva/multimodal/utils/image/__init__.py | 1 + .../eva/multimodal/utils/image/test_encode.py | 38 +++ tests/eva/multimodal/utils/text/__init__.py | 1 + .../multimodal/utils/text/test_messages.py | 65 +++++ .../data/transforms/spatial/test_resize.py | 58 +++++ 29 files changed, 1549 insertions(+) create mode 100644 tests/eva/language/data/dataloaders/collate_fn/test_text.py create mode 100644 tests/eva/language/data/datasets/test_text.py create mode 100644 tests/eva/language/models/wrappers/test_from_registry.py create mode 100644 tests/eva/language/utils/test_messages.py create mode 100644 tests/eva/multimodal/data/__init__.py create mode 100644 tests/eva/multimodal/data/dataloaders/__init__.py create mode 100644 tests/eva/multimodal/data/dataloaders/collate_fn/__init__.py create mode 100644 tests/eva/multimodal/data/dataloaders/collate_fn/test_text_image.py create mode 100644 tests/eva/multimodal/data/datasets/__init__.py create mode 100644 tests/eva/multimodal/data/datasets/multiple_choice/__init__.py create mode 100644 tests/eva/multimodal/data/datasets/multiple_choice/test_patch_camelyon.py create mode 100644 tests/eva/multimodal/data/datasets/test_base.py create mode 100644 tests/eva/multimodal/data/datasets/test_dataset_typings.py create mode 100644 tests/eva/multimodal/data/datasets/test_schemas.py create mode 100644 tests/eva/multimodal/data/datasets/test_text_image.py create mode 100644 tests/eva/multimodal/models/__init__.py create mode 100644 tests/eva/multimodal/models/modules/test_vision_language.py create mode 100644 tests/eva/multimodal/models/test_typings.py create mode 100644 tests/eva/multimodal/models/wrappers/__init__.py create mode 100644 tests/eva/multimodal/models/wrappers/test_base.py create mode 100644 tests/eva/multimodal/models/wrappers/test_from_registry.py create mode 100644 tests/eva/multimodal/models/wrappers/test_huggingface.py create mode 100644 tests/eva/multimodal/models/wrappers/test_litellm.py create mode 100644 tests/eva/multimodal/utils/__init__.py create mode 100644 tests/eva/multimodal/utils/image/__init__.py create mode 100644 tests/eva/multimodal/utils/image/test_encode.py create mode 100644 tests/eva/multimodal/utils/text/__init__.py create mode 100644 tests/eva/multimodal/utils/text/test_messages.py create mode 100644 tests/eva/vision/data/transforms/spatial/test_resize.py diff --git a/tests/eva/language/data/dataloaders/collate_fn/test_text.py b/tests/eva/language/data/dataloaders/collate_fn/test_text.py new file mode 100644 index 000000000..5aadc8e51 --- /dev/null +++ b/tests/eva/language/data/dataloaders/collate_fn/test_text.py @@ -0,0 +1,182 @@ +"""Tests for text collate functions.""" + +import torch + +from eva.language.data.dataloaders.collate_fn.text import text_collate +from eva.language.data.datasets.typings import TextSample +from eva.language.data.messages import UserMessage + + +def test_text_collate_with_targets(): + """Test collating samples with targets.""" + samples = [ + TextSample( + text=[UserMessage(content="Text 1")], + target=torch.tensor(0), + metadata={"id": 1, "category": "A"}, + ), + TextSample( + text=[UserMessage(content="Text 2")], + target=torch.tensor(1), + metadata={"id": 2, "category": "B"}, + ), + TextSample( + text=[UserMessage(content="Text 3")], + target=torch.tensor(2), + metadata={"id": 3, "category": "A"}, + ), + ] + + batch = text_collate(samples) + + assert len(batch.text) == 3 + assert batch.text[0][0].content == "Text 1" + assert batch.text[1][0].content == "Text 2" + assert batch.text[2][0].content == "Text 3" + assert batch.target is not None + assert batch.target.shape == (3,) + assert torch.equal(batch.target, torch.tensor([0, 1, 2])) + assert batch.metadata == {"id": [1, 2, 3], "category": ["A", "B", "A"]} + + +def test_text_collate_without_targets(): + """Test collating samples without targets.""" + samples = [ + TextSample( + text=[UserMessage(content="Text A")], + target=None, + metadata={"key": "val1", "score": 0.5}, + ), + TextSample( + text=[UserMessage(content="Text B")], + target=None, + metadata={"key": "val2", "score": 0.8}, + ), + ] + + batch = text_collate(samples) + + assert len(batch.text) == 2 + assert batch.text[0][0].content == "Text A" + assert batch.text[1][0].content == "Text B" + assert batch.target is None + assert batch.metadata == {"key": ["val1", "val2"], "score": [0.5, 0.8]} + + +def test_text_collate_without_metadata(): + """Test collating samples without metadata.""" + samples = [ + TextSample( + text=[UserMessage(content="Text 1")], + target=torch.tensor(0), + metadata=None, + ), + TextSample( + text=[UserMessage(content="Text 2")], + target=torch.tensor(1), + metadata=None, + ), + ] + + batch = text_collate(samples) + + assert len(batch.text) == 2 + assert batch.text[0][0].content == "Text 1" + assert batch.text[1][0].content == "Text 2" + assert batch.target is not None + assert batch.target.shape == (2,) + assert torch.equal(batch.target, torch.tensor([0, 1])) + assert batch.metadata is None + + +def test_text_collate_with_multiple_messages(): + """Test collating samples with multiple messages in conversation.""" + samples = [ + TextSample( + text=[ + UserMessage(content="Question 1"), + UserMessage(content="Follow-up 1"), + ], + target=torch.tensor([1, 0, 0]), + metadata={"sample_id": "s1"}, + ), + TextSample( + text=[ + UserMessage(content="Question 2"), + UserMessage(content="Follow-up 2"), + ], + target=torch.tensor([0, 1, 0]), + metadata={"sample_id": "s2"}, + ), + ] + + batch = text_collate(samples) + + assert len(batch.text) == 2 + assert len(batch.text[0]) == 2 + assert len(batch.text[1]) == 2 + assert batch.text[0][0].content == "Question 1" + assert batch.text[0][1].content == "Follow-up 1" + assert batch.text[1][0].content == "Question 2" + assert batch.text[1][1].content == "Follow-up 2" + assert batch.target is not None + assert batch.target.shape == (2, 3) + assert torch.equal(batch.target, torch.tensor([[1, 0, 0], [0, 1, 0]])) + assert batch.metadata == {"sample_id": ["s1", "s2"]} + + +def test_text_collate_with_mixed_metadata(): + """Test collating samples where some have metadata and some don't.""" + samples = [ + TextSample( + text=[UserMessage(content="Text with metadata")], + target=torch.tensor(0.5), + metadata={"has_meta": True}, + ), + TextSample( + text=[UserMessage(content="Text without metadata")], + target=torch.tensor(0.7), + metadata=None, + ), + ] + + batch = text_collate(samples) + + assert len(batch.text) == 2 + assert batch.text[0][0].content == "Text with metadata" + assert batch.text[1][0].content == "Text without metadata" + assert batch.target is not None + assert batch.target.shape == (2,) + assert torch.allclose(batch.target, torch.tensor([0.5, 0.7])) + assert batch.metadata == {"has_meta": [True]} + + +def test_text_collate_empty_batch(): + """Test collating an empty batch.""" + samples = [] + + try: + _ = text_collate(samples) + raise AssertionError("Should raise an error for empty batch") + except (ValueError, IndexError): + pass + + +def test_text_collate_single_sample(): + """Test collating a single sample.""" + samples = [ + TextSample( + text=[UserMessage(content="Single text")], + target=torch.tensor([1.0, 2.0, 3.0]), + metadata={"single": True}, + ) + ] + + batch = text_collate(samples) + + assert len(batch.text) == 1 + assert batch.text[0][0].content == "Single text" + assert batch.target is not None + assert batch.target.shape == (1, 3) + assert torch.equal(batch.target, torch.tensor([[1.0, 2.0, 3.0]])) + assert batch.metadata == {"single": [True]} diff --git a/tests/eva/language/data/datasets/test_text.py b/tests/eva/language/data/datasets/test_text.py new file mode 100644 index 000000000..ad338fcbd --- /dev/null +++ b/tests/eva/language/data/datasets/test_text.py @@ -0,0 +1,80 @@ +"""Tests for TextDataset class.""" + +from typing import Any, Dict + +from typing_extensions import override + +from eva.language.data.datasets.schemas import TransformsSchema +from eva.language.data.datasets.text import TextDataset +from eva.language.data.datasets.typings import TextSample +from eva.language.data.messages import UserMessage + + +class ConcreteTextDataset(TextDataset): + """Concrete implementation for testing.""" + + def __init__(self, transforms: TransformsSchema | None = None): + """Initialize test dataset.""" + super().__init__(transforms=transforms) + self._size = 3 + + @override + def __len__(self) -> int: + return self._size + + @override + def load_text(self, index: int): + return [UserMessage(content=f"Text {index}")] + + @override + def load_target(self, index: int) -> int: + return index + + @override + def load_metadata(self, index: int) -> Dict[str, Any]: + return {"index": index} + + +def test_text_dataset_getitem(): + """Test __getitem__ returns proper TextSample.""" + dataset = ConcreteTextDataset() + sample = dataset[0] + + assert isinstance(sample, TextSample) + assert len(sample.text) == 1 + assert sample.text[0].content == "Text 0" + assert sample.target == 0 + assert sample.metadata == {"index": 0} + + +def test_text_dataset_with_transforms(): + """Test dataset applies transforms correctly.""" + + def text_transform(text): + # Modify the text content + return [UserMessage(content=text[0].content.upper())] + + def target_transform(target): + return target * 10 + + transforms = TransformsSchema( + text=text_transform, + target=target_transform, + ) + + dataset = ConcreteTextDataset(transforms=transforms) + sample = dataset[1] + + assert sample.text[0].content == "TEXT 1" + assert sample.target == 10 + assert sample.metadata == {"index": 1} + + +def test_text_dataset_without_transforms(): + """Test dataset without transforms returns original data.""" + dataset = ConcreteTextDataset(transforms=None) + sample = dataset[2] + + assert sample.text[0].content == "Text 2" + assert sample.target == 2 + assert sample.metadata == {"index": 2} diff --git a/tests/eva/language/models/wrappers/test_from_registry.py b/tests/eva/language/models/wrappers/test_from_registry.py new file mode 100644 index 000000000..a02f9232e --- /dev/null +++ b/tests/eva/language/models/wrappers/test_from_registry.py @@ -0,0 +1,28 @@ +"""Tests for language from registry model wrapper.""" + +import os +from unittest import mock + +import pytest + +from eva.language.models import wrappers + + +@pytest.mark.parametrize( + ("model_name", "model_class"), + [ + ("anthropic/claude-3-7-sonnet-20250219", wrappers.LiteLLMModel), + ], +) +def test_load_model(model_name: str, model_class: type): + """Test loading a model from the registry.""" + with mock.patch.dict( + os.environ, + { + "ANTHROPIC_API_KEY": "test_key", + }, + ): + model = wrappers.ModelFromRegistry(model_name) + + assert isinstance(model, wrappers.ModelFromRegistry) + assert isinstance(model.model, model_class) diff --git a/tests/eva/language/utils/test_messages.py b/tests/eva/language/utils/test_messages.py new file mode 100644 index 000000000..2f4d6859a --- /dev/null +++ b/tests/eva/language/utils/test_messages.py @@ -0,0 +1,241 @@ +"""Tests for message formatting utilities.""" + +from eva.language.data.messages import ( + AssistantMessage, + MessageSeries, + ModelSystemMessage, + SystemMessage, + TaskSystemMessage, + UserMessage, +) +from eva.language.utils.text.messages import ( + batch_insert_system_message, + combine_system_messages, + format_chat_message, + insert_system_message, + merge_message_contents, +) + + +def test_format_chat_message_single_message(): + """Test formatting a single message.""" + messages: MessageSeries = [UserMessage(content="Hello")] + formatted = format_chat_message(messages) + + assert len(formatted) == 1 + assert formatted[0]["role"] == "user" + assert formatted[0]["content"] == "Hello" + + +def test_format_chat_message_multiple_messages(): + """Test formatting multiple messages of different types.""" + messages: MessageSeries = [ + SystemMessage(content="System prompt"), + UserMessage(content="User question"), + AssistantMessage(content="Assistant response"), + ] + formatted = format_chat_message(messages) + + assert len(formatted) == 3 + assert formatted[0] == {"role": "system", "content": "System prompt"} + assert formatted[1] == {"role": "user", "content": "User question"} + assert formatted[2] == {"role": "assistant", "content": "Assistant response"} + + +def test_format_chat_message_empty(): + """Test formatting an empty message series.""" + messages: MessageSeries = [] + formatted = format_chat_message(messages) + + assert formatted == [] + + +def test_combine_system_messages_single_system(): + """Test combining a single system message (should remain unchanged).""" + messages: MessageSeries = [ + SystemMessage(content="System prompt"), + UserMessage(content="User question"), + ] + combined = combine_system_messages(messages) + + assert len(combined) == 2 + assert combined[0].role == "system" + assert combined[0].content == "System prompt" + assert combined[1].role == "user" + assert combined[1].content == "User question" + + +def test_combine_system_messages_multiple_system(): + """Test combining multiple system messages.""" + messages: MessageSeries = [ + SystemMessage(content="First system"), + ModelSystemMessage(content="Model instructions"), + TaskSystemMessage(content="Task instructions"), + UserMessage(content="User question"), + ] + combined = combine_system_messages(messages) + + assert len(combined) == 2 + assert combined[0].role == "system" + assert combined[0].content == "First system\nModel instructions\nTask instructions" + assert combined[1].role == "user" + assert combined[1].content == "User question" + + +def test_combine_system_messages_custom_join_char(): + """Test combining system messages with custom join character.""" + messages: MessageSeries = [ + SystemMessage(content="First"), + SystemMessage(content="Second"), + UserMessage(content="User"), + ] + combined = combine_system_messages(messages, join_char=" | ") + + assert len(combined) == 2 + assert combined[0].content == "First | Second" + assert combined[1].content == "User" + + +def test_combine_system_messages_no_system(): + """Test combining when there are no system messages.""" + messages: MessageSeries = [ + UserMessage(content="User question"), + AssistantMessage(content="Assistant response"), + ] + combined = combine_system_messages(messages) + + assert combined == messages + + +def test_combine_system_messages_only_system(): + """Test combining when there are only system messages.""" + messages: MessageSeries = [ + SystemMessage(content="First"), + ModelSystemMessage(content="Second"), + TaskSystemMessage(content="Third"), + ] + combined = combine_system_messages(messages) + + assert len(combined) == 1 + assert combined[0].role == "system" + assert combined[0].content == "First\nSecond\nThird" + + +def test_merge_message_contents_single(): + """Test merging contents of a single message.""" + messages: MessageSeries = [UserMessage(content="Hello")] + merged = merge_message_contents(messages) + + assert merged == "Hello" + + +def test_merge_message_contents_multiple(): + """Test merging contents of multiple messages.""" + messages: MessageSeries = [ + SystemMessage(content="System"), + UserMessage(content="User"), + AssistantMessage(content="Assistant"), + ] + merged = merge_message_contents(messages) + + assert merged == "System\nUser\nAssistant" + + +def test_merge_message_contents_custom_join_char(): + """Test merging contents with custom join character.""" + messages: MessageSeries = [ + UserMessage(content="First"), + UserMessage(content="Second"), + UserMessage(content="Third"), + ] + merged = merge_message_contents(messages, join_char=" -> ") + + assert merged == "First -> Second -> Third" + + +def test_merge_message_contents_empty(): + """Test merging empty message series.""" + messages: MessageSeries = [] + merged = merge_message_contents(messages) + + assert merged == "" + + +def test_insert_system_message_with_message(): + """Test inserting a system message.""" + messages: MessageSeries = [ + UserMessage(content="User question"), + AssistantMessage(content="Response"), + ] + system_msg = SystemMessage(content="System prompt") + result = insert_system_message(messages, system_msg) + + assert len(result) == 3 + assert result[0] == system_msg + assert result[1].content == "User question" + assert result[2].content == "Response" + + +def test_insert_system_message_none(): + """Test inserting None system message (should return original).""" + messages: MessageSeries = [ + UserMessage(content="User question"), + AssistantMessage(content="Response"), + ] + result = insert_system_message(messages, None) + + assert result == messages + + +def test_insert_system_message_empty_list(): + """Test inserting system message into empty list.""" + messages: MessageSeries = [] + system_msg = SystemMessage(content="System prompt") + result = insert_system_message(messages, system_msg) + + assert len(result) == 1 + assert result[0] == system_msg + + +def test_batch_insert_system_message(): + """Test inserting system message into multiple message series.""" + batch_messages = [ + [UserMessage(content="First user"), AssistantMessage(content="First assistant")], + [UserMessage(content="Second user")], + [], + ] + system_msg = SystemMessage(content="System prompt") + result = batch_insert_system_message(batch_messages, system_msg) + + assert len(result) == 3 + assert len(result[0]) == 3 + assert result[0][0] == system_msg + assert result[0][1].content == "First user" + assert result[0][2].content == "First assistant" + + assert len(result[1]) == 2 + assert result[1][0] == system_msg + assert result[1][1].content == "Second user" + + assert len(result[2]) == 1 + assert result[2][0] == system_msg + + +def test_batch_insert_system_message_none(): + """Test batch inserting None system message.""" + batch_messages = [ + [UserMessage(content="First user")], + [UserMessage(content="Second user")], + ] + result = batch_insert_system_message(batch_messages, None) # type: ignore + + assert result == batch_messages + + +def test_batch_insert_system_message_empty_batch(): + """Test batch inserting into empty batch.""" + batch_messages = [] + system_msg = SystemMessage(content="System prompt") + result = batch_insert_system_message(batch_messages, system_msg) + + assert result == [] diff --git a/tests/eva/multimodal/data/__init__.py b/tests/eva/multimodal/data/__init__.py new file mode 100644 index 000000000..a2bb357d6 --- /dev/null +++ b/tests/eva/multimodal/data/__init__.py @@ -0,0 +1 @@ +"""Test data utilities for multimodal models.""" diff --git a/tests/eva/multimodal/data/dataloaders/__init__.py b/tests/eva/multimodal/data/dataloaders/__init__.py new file mode 100644 index 000000000..b27ccf60d --- /dev/null +++ b/tests/eva/multimodal/data/dataloaders/__init__.py @@ -0,0 +1 @@ +"""Test dataloaders for multimodal models.""" diff --git a/tests/eva/multimodal/data/dataloaders/collate_fn/__init__.py b/tests/eva/multimodal/data/dataloaders/collate_fn/__init__.py new file mode 100644 index 000000000..2f2e4428a --- /dev/null +++ b/tests/eva/multimodal/data/dataloaders/collate_fn/__init__.py @@ -0,0 +1 @@ +"""Test collate functions for multimodal dataloaders.""" diff --git a/tests/eva/multimodal/data/dataloaders/collate_fn/test_text_image.py b/tests/eva/multimodal/data/dataloaders/collate_fn/test_text_image.py new file mode 100644 index 000000000..81974afbc --- /dev/null +++ b/tests/eva/multimodal/data/dataloaders/collate_fn/test_text_image.py @@ -0,0 +1,88 @@ +"""Tests for collate functions.""" + +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import UserMessage +from eva.multimodal.data.dataloaders.collate_fn.text_image import text_image_collate +from eva.multimodal.data.datasets.typings import TextImageSample + + +def test_text_image_collate_with_targets(): + """Test collating samples with targets.""" + samples = [ + TextImageSample( + text=[UserMessage(content="Text 1")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=torch.tensor(0), + metadata={"id": 1}, + ), + TextImageSample( + text=[UserMessage(content="Text 2")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=torch.tensor(1), + metadata={"id": 2}, + ), + ] + + batch = text_image_collate(samples) + + assert len(batch.text) == 2 + assert batch.text[0][0].content == "Text 1" + assert batch.text[1][0].content == "Text 2" + assert len(batch.image) == 2 + assert batch.target is not None + assert batch.target.shape == (2,) + assert torch.equal(batch.target, torch.tensor([0, 1])) + assert batch.metadata == {"id": [1, 2]} + + +def test_text_image_collate_without_targets(): + """Test collating samples without targets.""" + samples = [ + TextImageSample( + text=[UserMessage(content="Text A")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=None, + metadata={"key": "val1"}, + ), + TextImageSample( + text=[UserMessage(content="Text B")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=None, + metadata={"key": "val2"}, + ), + ] + + batch = text_image_collate(samples) + + assert len(batch.text) == 2 + assert len(batch.image) == 2 + assert batch.target is None + assert batch.metadata == {"key": ["val1", "val2"]} + + +def test_text_image_collate_without_metadata(): + """Test collating samples without metadata.""" + samples = [ + TextImageSample( + text=[UserMessage(content="Text")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=torch.tensor(0), + metadata=None, + ), + TextImageSample( + text=[UserMessage(content="Text")], + image=tv_tensors.Image(torch.rand(3, 224, 224)), + target=torch.tensor(1), + metadata=None, + ), + ] + + batch = text_image_collate(samples) + + assert len(batch.text) == 2 + assert len(batch.image) == 2 + assert batch.target is not None + assert batch.target.shape == (2,) + assert batch.metadata is None diff --git a/tests/eva/multimodal/data/datasets/__init__.py b/tests/eva/multimodal/data/datasets/__init__.py new file mode 100644 index 000000000..7f1a64830 --- /dev/null +++ b/tests/eva/multimodal/data/datasets/__init__.py @@ -0,0 +1 @@ +"""Test datasets for multimodal models.""" diff --git a/tests/eva/multimodal/data/datasets/multiple_choice/__init__.py b/tests/eva/multimodal/data/datasets/multiple_choice/__init__.py new file mode 100644 index 000000000..3ba318ec8 --- /dev/null +++ b/tests/eva/multimodal/data/datasets/multiple_choice/__init__.py @@ -0,0 +1 @@ +"""Tests for multiple choice multimodal datasets.""" diff --git a/tests/eva/multimodal/data/datasets/multiple_choice/test_patch_camelyon.py b/tests/eva/multimodal/data/datasets/multiple_choice/test_patch_camelyon.py new file mode 100644 index 000000000..1539e450c --- /dev/null +++ b/tests/eva/multimodal/data/datasets/multiple_choice/test_patch_camelyon.py @@ -0,0 +1,104 @@ +"""PatchCamelyon multimodal dataset tests.""" + +import os +from typing import Literal + +import pytest +from torchvision import tv_tensors + +from eva.language.data.messages import Message, UserMessage +from eva.multimodal.data.datasets.multiple_choice import patch_camelyon +from eva.multimodal.data.datasets.typings import TextImageSample + + +@pytest.mark.parametrize( + "split, expected_length", + [("train", 4), ("val", 2), ("test", 1)], +) +def test_length(patch_camelyon_dataset: patch_camelyon.PatchCamelyon, expected_length: int) -> None: + """Tests the length of the dataset.""" + assert len(patch_camelyon_dataset) == expected_length + + +@pytest.mark.parametrize( + "split", + ["train", "val", "test"], +) +def test_sample(patch_camelyon_dataset: patch_camelyon.PatchCamelyon) -> None: + """Tests the format of a dataset sample.""" + sample = patch_camelyon_dataset[0] + # assert data sample is a TextImageSample + assert isinstance(sample, TextImageSample) + + # Test text component + assert isinstance(sample.text, list) + assert len(sample.text) == 1 + assert isinstance(sample.text[0], UserMessage) + assert "metastatic breast tissue" in sample.text[0].content + + # Test image component + assert isinstance(sample.image, tv_tensors.Image) + assert sample.image.shape == (3, 96, 96) + + # Test target + assert isinstance(sample.target, int) + assert sample.target in [0, 1] + + # Test metadata + assert sample.metadata is not None + + +@pytest.mark.parametrize( + "split", + ["train", "val", "test"], +) +def test_custom_prompt(split: Literal["train", "val", "test"], assets_path: str) -> None: + """Tests the dataset with a custom prompt.""" + custom_prompt = "Is this image showing cancer? Answer: " + dataset = patch_camelyon.PatchCamelyon( + root=os.path.join(assets_path, "vision", "datasets", "patch_camelyon"), + split=split, + prompt=custom_prompt, + ) + + sample = dataset[0] + assert isinstance(sample.text, list) + assert isinstance(sample.text[0], Message) + assert sample.text[0].content == custom_prompt + + +@pytest.mark.parametrize( + "split", + ["train", "val", "test"], +) +def test_max_samples(split: Literal["train", "val", "test"], assets_path: str) -> None: + """Tests the dataset with max_samples limit.""" + max_samples = 1 + dataset = patch_camelyon.PatchCamelyon( + root=os.path.join(assets_path, "vision", "datasets", "patch_camelyon"), + split=split, + max_samples=max_samples, + ) + + assert len(dataset) == max_samples + + +def test_class_to_idx(assets_path: str) -> None: + """Tests the class_to_idx mapping.""" + dataset = patch_camelyon.PatchCamelyon( + root=os.path.join(assets_path, "vision", "datasets", "patch_camelyon"), + split="train", + ) + assert dataset.class_to_idx == {"A": 0, "B": 1} + + +@pytest.fixture(scope="function") +def patch_camelyon_dataset( + split: Literal["train", "val", "test"], assets_path: str +) -> patch_camelyon.PatchCamelyon: + """PatchCamelyon multimodal dataset fixture.""" + dataset = patch_camelyon.PatchCamelyon( + root=os.path.join(assets_path, "vision", "datasets", "patch_camelyon"), + split=split, + ) + return dataset diff --git a/tests/eva/multimodal/data/datasets/test_base.py b/tests/eva/multimodal/data/datasets/test_base.py new file mode 100644 index 000000000..faa7ac9ec --- /dev/null +++ b/tests/eva/multimodal/data/datasets/test_base.py @@ -0,0 +1,31 @@ +"""Tests for MultimodalDataset base class.""" + +from typing_extensions import override + +from eva.multimodal.data.datasets.base import MultimodalDataset + + +class ConcreteMultimodalDataset(MultimodalDataset): + """Concrete implementation for testing.""" + + def __init__(self): + """Initialize test dataset.""" + self._data = ["sample1", "sample2", "sample3"] + + @override + def __len__(self): + return len(self._data) + + @override + def __getitem__(self, index): + return self._data[index] + + +def test_multimodal_dataset_inheritance(): + """Test that MultimodalDataset can be inherited and used.""" + dataset = ConcreteMultimodalDataset() + + assert len(dataset) == 3 + assert dataset[0] == "sample1" + assert dataset[1] == "sample2" + assert dataset[2] == "sample3" diff --git a/tests/eva/multimodal/data/datasets/test_dataset_typings.py b/tests/eva/multimodal/data/datasets/test_dataset_typings.py new file mode 100644 index 000000000..36e2b846e --- /dev/null +++ b/tests/eva/multimodal/data/datasets/test_dataset_typings.py @@ -0,0 +1,50 @@ +"""Tests for multimodal dataset typings.""" + +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import MessageSeries, UserMessage +from eva.multimodal.data.datasets.typings import TextImageSample + + +def test_text_image_sample_creation(): + """Test TextImageSample creation and field access.""" + text: MessageSeries = [UserMessage(content="Test message")] + image = tv_tensors.Image(torch.rand(3, 224, 224)) + target = 1 + metadata = {"key": "value"} + + sample = TextImageSample(text=text, image=image, target=target, metadata=metadata) + + assert sample.text == text + assert sample.image is image + assert sample.target == target + assert sample.metadata == metadata + + +def test_text_image_sample_with_none_fields(): + """Test TextImageSample with None target and metadata.""" + text: MessageSeries = [UserMessage(content="Test")] + image = tv_tensors.Image(torch.rand(3, 224, 224)) + + sample = TextImageSample(text=text, image=image, target=None, metadata=None) + + assert sample.text == text + assert sample.image is image + assert sample.target is None + assert sample.metadata is None + + +def test_text_image_sample_unpacking(): + """Test TextImageSample can be unpacked.""" + text: MessageSeries = [UserMessage(content="Test")] + image = tv_tensors.Image(torch.rand(3, 224, 224)) + + sample = TextImageSample(text=text, image=image, target=42, metadata={"test": True}) + + unpacked_text, unpacked_image, unpacked_target, unpacked_metadata = sample + + assert unpacked_text == text + assert unpacked_image is image + assert unpacked_target == 42 + assert unpacked_metadata == {"test": True} diff --git a/tests/eva/multimodal/data/datasets/test_schemas.py b/tests/eva/multimodal/data/datasets/test_schemas.py new file mode 100644 index 000000000..dcc048933 --- /dev/null +++ b/tests/eva/multimodal/data/datasets/test_schemas.py @@ -0,0 +1,48 @@ +"""Tests for multimodal dataset schemas.""" + +from eva.multimodal.data.datasets.schemas import TransformsSchema + + +def test_transforms_schema_with_all_fields(): + """Test TransformsSchema with all transform fields.""" + + def text_transform(x): + return x + + def image_transform(x): + return x + + def target_transform(x): + return x + + schema = TransformsSchema(text=text_transform, image=image_transform, target=target_transform) + + assert schema.text is text_transform + assert schema.image is image_transform + assert schema.target is target_transform + + +def test_transforms_schema_with_image_only(): + """Test TransformsSchema with only image transform.""" + + def image_transform(x): + return x + + schema = TransformsSchema(image=image_transform) + + assert schema.text is None + assert schema.image is image_transform + assert schema.target is None + + +def test_transforms_schema_frozen(): + """Test that TransformsSchema is frozen (immutable).""" + schema = TransformsSchema() + + # Attempt to modify should raise an error + try: + schema.image = lambda x: x # type: ignore + raise AssertionError("Schema should be frozen") + except Exception: + # Expected - schema is frozen + pass diff --git a/tests/eva/multimodal/data/datasets/test_text_image.py b/tests/eva/multimodal/data/datasets/test_text_image.py new file mode 100644 index 000000000..8857981db --- /dev/null +++ b/tests/eva/multimodal/data/datasets/test_text_image.py @@ -0,0 +1,90 @@ +"""Tests for TextImageDataset class.""" + +from typing import Any, Dict + +import torch +from torchvision import tv_tensors +from typing_extensions import override + +from eva.language.data.messages import UserMessage +from eva.multimodal.data.datasets.schemas import TransformsSchema +from eva.multimodal.data.datasets.text_image import TextImageDataset +from eva.multimodal.data.datasets.typings import TextImageSample + + +class ConcreteTextImageDataset(TextImageDataset): + """Concrete implementation for testing.""" + + def __init__(self, transforms: TransformsSchema | None = None): + """Initialize test dataset.""" + super().__init__(transforms=transforms) + self._size = 3 + + @override + def __len__(self) -> int: + return self._size + + @override + def load_text(self, index: int): + return [UserMessage(content=f"Text {index}")] + + @override + def load_image(self, index: int) -> tv_tensors.Image: + return tv_tensors.Image(torch.rand(3, 224, 224)) + + @override + def load_target(self, index: int) -> int: + return index + + @override + def load_metadata(self, index: int) -> Dict[str, Any]: + return {"index": index} + + +def test_text_image_dataset_getitem(): + """Test __getitem__ returns proper TextImageSample.""" + dataset = ConcreteTextImageDataset() + sample = dataset[0] + + assert isinstance(sample, TextImageSample) + assert len(sample.text) == 1 + assert sample.text[0].content == "Text 0" + assert isinstance(sample.image, tv_tensors.Image) + assert sample.target == 0 + assert sample.metadata == {"index": 0} + + +def test_text_image_dataset_with_transforms(): + """Test dataset applies transforms correctly.""" + + def text_transform(text): + # Modify the text content + return [UserMessage(content=text[0].content.upper())] + + def image_transform(image): + # Simple transform - just return a different sized image + return tv_tensors.Image(torch.rand(3, 128, 128)) + + def target_transform(target): + return target * 10 + + transforms = TransformsSchema( + text=text_transform, image=image_transform, target=target_transform + ) + + dataset = ConcreteTextImageDataset(transforms=transforms) + sample = dataset[1] + + assert sample.text[0].content == "TEXT 1" + assert sample.image.shape == (3, 128, 128) + assert sample.target == 10 + + +def test_text_image_dataset_without_transforms(): + """Test dataset without transforms returns original data.""" + dataset = ConcreteTextImageDataset(transforms=None) + sample = dataset[2] + + assert sample.text[0].content == "Text 2" + assert sample.image.shape == (3, 224, 224) + assert sample.target == 2 diff --git a/tests/eva/multimodal/models/__init__.py b/tests/eva/multimodal/models/__init__.py new file mode 100644 index 000000000..4aa2e3933 --- /dev/null +++ b/tests/eva/multimodal/models/__init__.py @@ -0,0 +1 @@ +"""Test models for multimodal functionality.""" diff --git a/tests/eva/multimodal/models/modules/test_vision_language.py b/tests/eva/multimodal/models/modules/test_vision_language.py new file mode 100644 index 000000000..51f9fd251 --- /dev/null +++ b/tests/eva/multimodal/models/modules/test_vision_language.py @@ -0,0 +1,113 @@ +"""Tests the vision-language model module.""" + +import pytest +import torch +from torch import nn +from torchvision import tv_tensors + +from eva.language.data.messages import MessageSeries, UserMessage +from eva.multimodal.models.modules.vision_language import VisionLanguageModule +from eva.multimodal.models.typings import TextImageBatch + + +def test_forward(vision_language_module): + """Test the forward method of the VisionLanguageModule class.""" + text: list[MessageSeries] = [[UserMessage(content="Hello world")]] + batch = TextImageBatch( + text=text, + image=[tv_tensors.Image(torch.rand(3, 224, 224))], + target=torch.tensor([0]), + metadata={"id": [1]}, + ) + expected = ["Dummy response Nr. 0"] + result = vision_language_module.forward(batch) + assert result == expected + + +def test_validation_step(vision_language_module): + """Test the validation_step method of the VisionLanguageModule class.""" + text: list[MessageSeries] = [[UserMessage(content="What is in the image?")]] + images = [tv_tensors.Image(torch.rand(3, 224, 224))] + targets = torch.tensor([1]) + metadata = {"id": [1], "category": ["test"]} + batch = TextImageBatch(text=text, image=images, target=targets, metadata=metadata) + + output = vision_language_module.validation_step(batch) + + assert "inputs" in output + assert "predictions" in output + assert "targets" in output + assert "metadata" in output + assert output["inputs"] == text + assert output["predictions"] == ["Dummy response Nr. 0"] + assert torch.equal(output["targets"], targets) + assert output["metadata"] == metadata + + +def test_test_step(vision_language_module): + """Test the test_step method of the VisionLanguageModule class.""" + text: list[MessageSeries] = [ + [UserMessage(content="Describe this image")], + [UserMessage(content="What do you see?")], + ] + images = [ + tv_tensors.Image(torch.rand(3, 224, 224)), + tv_tensors.Image(torch.rand(3, 224, 224)), + ] + targets = torch.tensor([0, 1]) + metadata = {"id": [1, 2]} + batch = TextImageBatch(text=text, image=images, target=targets, metadata=metadata) + + output = vision_language_module.test_step(batch) + + assert "inputs" in output + assert "predictions" in output + assert "targets" in output + assert "metadata" in output + assert output["inputs"] == text + assert output["predictions"] == ["Dummy response Nr. 0", "Dummy response Nr. 1"] + assert torch.equal(output["targets"], targets) + assert output["metadata"] == metadata + + +def test_init_attributes(model): + """Test the attributes of the VisionLanguageModule class.""" + module_instance = VisionLanguageModule(model=model) + assert module_instance.model is model + assert module_instance.metrics is not None # MetricModule is created by default + assert module_instance._postprocess is not None # BatchPostProcess is created by default + + +def test_batch_step_without_targets(vision_language_module): + """Test the _batch_step method with None targets.""" + text: list[MessageSeries] = [[UserMessage(content="Test message")]] + images = [tv_tensors.Image(torch.rand(3, 224, 224))] + batch = TextImageBatch(text=text, image=images, target=None, metadata=None) + + output = vision_language_module.validation_step(batch) + + assert output["targets"] is None + assert output["metadata"] is None + assert output["inputs"] == text + assert output["predictions"] == ["Dummy response Nr. 0"] + + +class DummyVisionLanguageModel(nn.Module): + """A simple vision-language model for testing purposes.""" + + def forward(self, batch: TextImageBatch) -> list[str]: + """Generate text responses based on the batch size.""" + text, images, _, _ = batch + return [f"Dummy response Nr. {i}" for i in range(len(text))] + + +@pytest.fixture +def model(): + """Return a dummy model instance.""" + return DummyVisionLanguageModel() + + +@pytest.fixture +def vision_language_module(model): + """Return a VisionLanguageModule instance.""" + return VisionLanguageModule(model=model) diff --git a/tests/eva/multimodal/models/test_typings.py b/tests/eva/multimodal/models/test_typings.py new file mode 100644 index 000000000..740acd1f6 --- /dev/null +++ b/tests/eva/multimodal/models/test_typings.py @@ -0,0 +1,36 @@ +"""Tests for multimodal model typings.""" + +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import MessageSeries, UserMessage +from eva.multimodal.models.typings import TextImageBatch + + +def test_text_image_batch_creation(): + """Test TextImageBatch creation and field access.""" + messages: list[MessageSeries] = [[UserMessage(content="Test")]] + images = [tv_tensors.Image(torch.rand(3, 224, 224))] + target = torch.tensor([1]) + metadata = {"key": "value"} + + batch = TextImageBatch(text=messages, image=images, target=target, metadata=metadata) + + assert batch.text == messages + assert batch.image == images + assert batch.target is not None and torch.equal(batch.target, target) + assert batch.metadata == metadata + + +def test_text_image_batch_unpacking(): + """Test TextImageBatch can be unpacked.""" + messages: list[MessageSeries] = [[UserMessage(content="Test")]] + images = [tv_tensors.Image(torch.rand(3, 224, 224))] + + batch = TextImageBatch(text=messages, image=images, target=None, metadata=None) + + text, image, target, metadata = batch + assert text == messages + assert image == images + assert target is None + assert metadata is None diff --git a/tests/eva/multimodal/models/wrappers/__init__.py b/tests/eva/multimodal/models/wrappers/__init__.py new file mode 100644 index 000000000..8b1c7f23f --- /dev/null +++ b/tests/eva/multimodal/models/wrappers/__init__.py @@ -0,0 +1 @@ +"""Test wrappers for multimodal models.""" diff --git a/tests/eva/multimodal/models/wrappers/test_base.py b/tests/eva/multimodal/models/wrappers/test_base.py new file mode 100644 index 000000000..ae503c44d --- /dev/null +++ b/tests/eva/multimodal/models/wrappers/test_base.py @@ -0,0 +1,46 @@ +"""Tests for VisionLanguageModel base class.""" + +from typing import Any, List +from unittest.mock import MagicMock + +from typing_extensions import override + +from eva.multimodal.models.typings import TextImageBatch +from eva.multimodal.models.wrappers.base import VisionLanguageModel + + +class ConcreteVisionLanguageModel(VisionLanguageModel): + """Concrete implementation for testing.""" + + def __init__(self, system_prompt: str | None = None): + """Initialize test model.""" + super().__init__(system_prompt=system_prompt) + self.model = MagicMock() + + @override + def format_inputs(self, batch: TextImageBatch) -> Any: + return {"formatted": batch} + + @override + def model_forward(self, batch: Any) -> List[str]: + return ["response1", "response2"] + + +def test_system_prompt_initialization(): + """Test that system prompt is correctly initialized.""" + model_with_prompt = ConcreteVisionLanguageModel(system_prompt="You are a helpful assistant") + assert model_with_prompt.system_message is not None + assert model_with_prompt.system_message.content == "You are a helpful assistant" + + model_without_prompt = ConcreteVisionLanguageModel(system_prompt=None) + assert model_without_prompt.system_message is None + + +def test_forward_delegates_to_format_and_model_forward(): + """Test that forward correctly delegates to format_inputs and model_forward.""" + model = ConcreteVisionLanguageModel() + batch = MagicMock(spec=TextImageBatch) + + result = model.forward(batch) + + assert result == ["response1", "response2"] diff --git a/tests/eva/multimodal/models/wrappers/test_from_registry.py b/tests/eva/multimodal/models/wrappers/test_from_registry.py new file mode 100644 index 000000000..1d7ab3758 --- /dev/null +++ b/tests/eva/multimodal/models/wrappers/test_from_registry.py @@ -0,0 +1,28 @@ +"""Tests for multimodal from registry model wrapper.""" + +import os +from unittest import mock + +import pytest + +from eva.multimodal.models import wrappers + + +@pytest.mark.parametrize( + ("model_name", "model_class"), + [ + ("anthropic/claude-3-7-sonnet-20250219", wrappers.LiteLLMModel), + ], +) +def test_load_model(model_name: str, model_class: type): + """Test loading a model from the registry.""" + with mock.patch.dict( + os.environ, + { + "ANTHROPIC_API_KEY": "test_key", + }, + ): + model = wrappers.ModelFromRegistry(model_name) + + assert isinstance(model, wrappers.ModelFromRegistry) + assert isinstance(model.model, model_class) diff --git a/tests/eva/multimodal/models/wrappers/test_huggingface.py b/tests/eva/multimodal/models/wrappers/test_huggingface.py new file mode 100644 index 000000000..adf4e7f14 --- /dev/null +++ b/tests/eva/multimodal/models/wrappers/test_huggingface.py @@ -0,0 +1,124 @@ +"""HuggingFace multimodal wrapper tests.""" + +from unittest.mock import MagicMock, patch + +import pytest +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import UserMessage +from eva.multimodal.models.typings import TextImageBatch +from eva.multimodal.models.wrappers.huggingface import HuggingFaceModel + + +@pytest.mark.parametrize( + "model_name, model_class, with_image", + [ + ("llava-hf/llava-1.5-7b-hf", "LlavaForConditionalGeneration", True), + ("llava-hf/llava-1.5-7b-hf", "LlavaForConditionalGeneration", False), + ], +) +def test_huggingface_model_generation(model_name: str, model_class: str, with_image: bool): + """Test HuggingFace multimodal model generation with mocked components.""" + mock_processor = MagicMock() + mock_processor.chat_template = "template" + mock_processor.apply_chat_template.return_value = "formatted text" + mock_processor.return_value.to.return_value = {"input_ids": torch.tensor([[1, 2, 3]])} + mock_processor.batch_decode.side_effect = [[""], ["Generated response"]] + + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + mock_model.generate.return_value = torch.tensor([[1, 2, 3, 4, 5]]) + + with ( + patch("transformers.AutoProcessor.from_pretrained", return_value=mock_processor), + patch(f"transformers.{model_class}.from_pretrained", return_value=mock_model), + ): + model = HuggingFaceModel( + model_name_or_path=model_name, + model_class=model_class, + generation_kwargs={"max_new_tokens": 50}, + ) + + # Always create an image tensor, even if not used + image = tv_tensors.Image(torch.rand(3, 224, 224)) + batch = TextImageBatch( + text=[[UserMessage(content="Describe this")]], + image=[image], + target=None, + metadata={}, + ) + + result = model(batch) + assert result == ["Generated response"] + assert mock_model.generate.called + + +def test_format_inputs_with_image(): + """Test format_inputs correctly handles image inputs.""" + mock_processor = MagicMock() + mock_processor.chat_template = "template" + mock_processor.apply_chat_template.return_value = "formatted text" + mock_processor.return_value.to.return_value = { + "input_ids": torch.tensor([[1, 2, 3]]), + "pixel_values": torch.rand(1, 3, 224, 224), + } + + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + + with ( + patch("transformers.AutoProcessor.from_pretrained", return_value=mock_processor), + patch( + "transformers.LlavaForConditionalGeneration.from_pretrained", return_value=mock_model + ), + ): + model = HuggingFaceModel( + model_name_or_path="test-model", + model_class="LlavaForConditionalGeneration", + ) + + image = tv_tensors.Image(torch.rand(3, 224, 224)) + batch = TextImageBatch( + text=[[UserMessage(content="Test")]], + image=[image], + target=None, + metadata={}, + ) + + formatted = model.format_inputs(batch) + + mock_processor.assert_called_with( + text=["formatted text"], + image=[[image]], + return_tensors="pt", + ) + assert "input_ids" in formatted + + +def test_decode_output(): + """Test _decode_output correctly decodes model output.""" + mock_processor = MagicMock() + mock_processor.batch_decode.side_effect = [["Input text"], ["Output text"]] + + mock_model = MagicMock() + mock_model.device = torch.device("cpu") + + with ( + patch("transformers.AutoProcessor.from_pretrained", return_value=mock_processor), + patch( + "transformers.LlavaForConditionalGeneration.from_pretrained", return_value=mock_model + ), + ): + model = HuggingFaceModel( + model_name_or_path="test-model", + model_class="LlavaForConditionalGeneration", + ) + + output = torch.tensor([[1, 2, 3, 4, 5, 6]]) + instruction_length = 3 + + decoded = model._decode_output(output, instruction_length) + + assert decoded == ["Output text"] + assert mock_processor.batch_decode.call_count == 2 diff --git a/tests/eva/multimodal/models/wrappers/test_litellm.py b/tests/eva/multimodal/models/wrappers/test_litellm.py new file mode 100644 index 000000000..4637aeb49 --- /dev/null +++ b/tests/eva/multimodal/models/wrappers/test_litellm.py @@ -0,0 +1,89 @@ +"""LiteLLM multimodal wrapper tests.""" + +import pytest +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import UserMessage +from eva.multimodal.models.typings import TextImageBatch +from eva.multimodal.models.wrappers.litellm import LiteLLMModel + +DUMMY_RESPONSE = {"choices": [{"message": {"content": "Test response", "role": "assistant"}}]} + + +def test_generate_with_image(model_instance, sample_image): + """Test that the generate method works with image input.""" + batch = TextImageBatch( + text=[[UserMessage(content="What's in this image?")]], + image=[sample_image], + target=None, + metadata={}, + ) + result = model_instance(batch) + assert result == ["Test response"] + + +def test_generate_without_image(model_instance): + """Test that the generate method works without image input.""" + # Create a dummy image tensor, but we'll mock the response anyway + dummy_image = tv_tensors.Image(torch.zeros(3, 1, 1)) + batch = TextImageBatch( + text=[[UserMessage(content="Hello, world!")]], + image=[dummy_image], + target=None, + metadata={}, + ) + result = model_instance(batch) + assert result == ["Test response"] + + +def test_format_inputs_with_image(model_instance, sample_image): + """Test format_inputs properly formats messages with images.""" + batch = TextImageBatch( + text=[[UserMessage(content="Describe this")]], + image=[sample_image], + target=None, + metadata={}, + ) + formatted = model_instance.format_inputs(batch) + + assert isinstance(formatted, list) + assert len(formatted) == 1 + assert len(formatted[0]) == 2 # System message + user message + assert formatted[0][0]["role"] == "system" + assert formatted[0][0]["content"] == "You are a helpful assistant." + assert formatted[0][1]["role"] == "user" + assert isinstance(formatted[0][1]["content"], list) + assert len(formatted[0][1]["content"]) == 2 + assert formatted[0][1]["content"][0]["type"] == "text" + assert formatted[0][1]["content"][1]["type"] == "image_url" + + +@pytest.fixture +def sample_image(): + """Create a sample image tensor for testing.""" + return tv_tensors.Image(torch.rand(3, 224, 224)) + + +@pytest.fixture +def fake_completion(monkeypatch): + """Fixture to override `batch_completion` function and set a dummy OPENAI_API_KEY.""" + + def _fake_batch_completion(**_kwargs): + return [DUMMY_RESPONSE] + + monkeypatch.setenv("OPENAI_API_KEY", "dummy-key") + monkeypatch.setattr( + "eva.language.models.wrappers.litellm.batch_completion", _fake_batch_completion + ) + return _fake_batch_completion + + +@pytest.fixture +def model_instance(fake_completion): # noqa: ARG001 + """Fixture to instantiate the multimodal LiteLLMModel.""" + return LiteLLMModel( + "openai/gpt-4-vision-preview", + model_kwargs={"temperature": 0.7}, + system_prompt="You are a helpful assistant.", + ) diff --git a/tests/eva/multimodal/utils/__init__.py b/tests/eva/multimodal/utils/__init__.py new file mode 100644 index 000000000..0dbe6c22c --- /dev/null +++ b/tests/eva/multimodal/utils/__init__.py @@ -0,0 +1 @@ +"""Test utilities for multimodal functionality.""" diff --git a/tests/eva/multimodal/utils/image/__init__.py b/tests/eva/multimodal/utils/image/__init__.py new file mode 100644 index 000000000..e7f7758ec --- /dev/null +++ b/tests/eva/multimodal/utils/image/__init__.py @@ -0,0 +1 @@ +"""Test image utilities for multimodal models.""" diff --git a/tests/eva/multimodal/utils/image/test_encode.py b/tests/eva/multimodal/utils/image/test_encode.py new file mode 100644 index 000000000..3903c2272 --- /dev/null +++ b/tests/eva/multimodal/utils/image/test_encode.py @@ -0,0 +1,38 @@ +"""Tests for image encoding utilities.""" + +import base64 + +import pytest +import torch +from torchvision import tv_tensors + +from eva.multimodal.utils.image.encode import encode_image + + +def test_encode_image_base64(): + """Test base64 encoding of image tensors.""" + image = tv_tensors.Image(torch.rand(3, 224, 224)) + encoded = encode_image(image, encoding="base64") + + assert isinstance(encoded, str) + # Test that it's valid base64 + base64.b64decode(encoded) + assert len(encoded) > 0 + + +def test_encode_image_unsupported_encoding(): + """Test that unsupported encoding raises ValueError.""" + image = tv_tensors.Image(torch.rand(3, 224, 224)) + + with pytest.raises(ValueError, match="Unsupported encoding type"): + encode_image(image, encoding="unsupported") # type: ignore + + +@pytest.mark.parametrize("image_shape", [(3, 32, 32), (3, 224, 224), (3, 512, 512)]) +def test_encode_different_sizes(image_shape): + """Test encoding works with different image sizes.""" + image = tv_tensors.Image(torch.rand(*image_shape)) + encoded = encode_image(image, encoding="base64") + + assert isinstance(encoded, str) + assert len(encoded) > 0 diff --git a/tests/eva/multimodal/utils/text/__init__.py b/tests/eva/multimodal/utils/text/__init__.py new file mode 100644 index 000000000..376e1113a --- /dev/null +++ b/tests/eva/multimodal/utils/text/__init__.py @@ -0,0 +1 @@ +"""Test text utilities for multimodal models.""" diff --git a/tests/eva/multimodal/utils/text/test_messages.py b/tests/eva/multimodal/utils/text/test_messages.py new file mode 100644 index 000000000..7b726aab8 --- /dev/null +++ b/tests/eva/multimodal/utils/text/test_messages.py @@ -0,0 +1,65 @@ +"""Tests for message formatting utilities.""" + +import torch +from torchvision import tv_tensors + +from eva.language.data.messages import MessageSeries, SystemMessage, UserMessage +from eva.multimodal.utils.text.messages import format_huggingface_message, format_litellm_message + + +def test_format_huggingface_message_without_images(): + """Test formatting messages for HuggingFace without images.""" + messages: MessageSeries = [UserMessage(content="Hello")] + formatted = format_huggingface_message(messages, with_images=False) + + assert len(formatted) == 1 + assert formatted[0]["role"] == "user" + assert formatted[0]["content"] == "Hello" + + +def test_format_huggingface_message_with_images(): + """Test formatting messages for HuggingFace with images.""" + messages: MessageSeries = [ + SystemMessage(content="System prompt"), + UserMessage(content="What's this?"), + ] + formatted = format_huggingface_message(messages, with_images=True) + + assert len(formatted) == 2 + assert formatted[0]["role"] == "system" + assert formatted[0]["content"] == "System prompt" + assert formatted[1]["role"] == "user" + assert isinstance(formatted[1]["content"], list) + assert formatted[1]["content"][0]["type"] == "text" + assert formatted[1]["content"][1]["type"] == "image" + + +def test_format_litellm_message_without_image(): + """Test formatting messages for LiteLLM without image.""" + messages: MessageSeries = [UserMessage(content="Hello")] + formatted = format_litellm_message(messages, image=None) + + assert len(formatted) == 1 + assert formatted[0]["role"] == "user" + assert formatted[0]["content"] == "Hello" + + +def test_format_litellm_message_with_image(): + """Test formatting messages for LiteLLM with image.""" + messages: MessageSeries = [ + SystemMessage(content="System prompt"), + UserMessage(content="Describe this"), + ] + image = tv_tensors.Image(torch.rand(3, 224, 224)) + formatted = format_litellm_message(messages, image=image) + + assert len(formatted) == 2 + assert formatted[0]["role"] == "system" + assert formatted[0]["content"] == "System prompt" + assert formatted[1]["role"] == "user" + assert isinstance(formatted[1]["content"], list) + assert formatted[1]["content"][0]["type"] == "text" + assert formatted[1]["content"][0]["text"] == "Describe this" + assert formatted[1]["content"][1]["type"] == "image_url" + assert "url" in formatted[1]["content"][1]["image_url"] + assert formatted[1]["content"][1]["image_url"]["url"].startswith("data:image/png;base64,") diff --git a/tests/eva/vision/data/transforms/spatial/test_resize.py b/tests/eva/vision/data/transforms/spatial/test_resize.py new file mode 100644 index 000000000..ec4a1c812 --- /dev/null +++ b/tests/eva/vision/data/transforms/spatial/test_resize.py @@ -0,0 +1,58 @@ +"""Tests for resize transforms.""" + +import pytest +import torch +from torchvision import tv_tensors + +from eva.vision.data import transforms + + +def test_resize_with_size_only(): + """Test Resize with only size parameter provided.""" + resize_transform = transforms.Resize(size=(100, 100)) + test_image = tv_tensors.Image(torch.rand(3, 200, 200)) + + result = resize_transform(test_image) + + assert isinstance(result, tv_tensors.Image) + assert result.shape == (3, 100, 100) + + +def test_resize_with_max_bytes_only(): + """Test Resize with only max_bytes parameter provided.""" + resize_transform = transforms.Resize(max_bytes=1000) # Very small size to force resizing + test_image = tv_tensors.Image(torch.rand(3, 500, 500)) + + result = resize_transform(test_image) + + assert isinstance(result, tv_tensors.Image) + # Image should be smaller than original due to byte size constraint + assert result.shape[1] < 500 or result.shape[2] < 500 + + +def test_resize_with_both_size_and_max_bytes_raises_error(): + """Test Resize raises ValueError when both size and max_bytes parameters are provided.""" + with pytest.raises(ValueError, match="Cannot provide both 'size' and 'max_bytes' parameters."): + transforms.Resize(size=(150, 150), max_bytes=1000) + + +def test_resize_with_no_parameters(): + """Test Resize with neither size nor max_bytes parameters provided.""" + resize_transform = transforms.Resize() + test_image = tv_tensors.Image(torch.rand(3, 200, 200)) + + result = resize_transform(test_image) + + assert isinstance(result, tv_tensors.Image) + # Should return original image unchanged + assert result.shape == test_image.shape + assert torch.equal(result, test_image) + + +def test_resize_with_invalid_max_bytes(): + """Test Resize raises ValueError for non-positive max_bytes.""" + with pytest.raises(ValueError, match="'max_bytes' must be a positive integer."): + transforms.Resize(max_bytes=0) + + with pytest.raises(ValueError, match="'max_bytes' must be a positive integer."): + transforms.Resize(max_bytes=-100) From c66e0e0662b2ee79ab02e31706fb6708c296eb36 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicolas=20K=C3=A4nzig?= <36882833+nkaenzig@users.noreply.github.com> Date: Wed, 3 Sep 2025 10:25:19 +0200 Subject: [PATCH 3/3] Add support for `offline` mode to `language` and `multimodal` (#873) --- .gitattributes | 2 + .../offline/multiple_choice/pubmedqa.yaml | 74 ++++++ .../online/multiple_choice/pubmedqa.yaml | 16 +- .../multiple_choice/patch_camelyon.yaml | 84 ++++++ .../multiple_choice/patch_camelyon.yaml | 65 ----- src/eva/core/interface/interface.py | 21 ++ src/eva/core/models/modules/module.py | 4 +- src/eva/language/callbacks/__init__.py | 5 + .../language/callbacks/writers/__init__.py | 5 + .../language/callbacks/writers/prediction.py | 176 +++++++++++++ src/eva/language/data/dataloaders/__init__.py | 4 +- .../data/dataloaders/collate_fn/__init__.py | 4 +- .../data/dataloaders/collate_fn/text.py | 29 ++- src/eva/language/data/datasets/__init__.py | 2 + .../data/datasets/classification/pubmedqa.py | 28 +- src/eva/language/data/datasets/prediction.py | 151 +++++++++++ src/eva/language/data/datasets/schemas.py | 3 + src/eva/language/data/datasets/text.py | 1 - src/eva/language/data/datasets/typings.py | 16 ++ src/eva/language/data/messages.py | 15 +- src/eva/language/models/__init__.py | 4 +- src/eva/language/models/modules/__init__.py | 4 +- src/eva/language/models/modules/language.py | 40 ++- src/eva/language/models/typings.py | 16 ++ src/eva/language/models/wrappers/__init__.py | 11 +- src/eva/language/utils/str_to_int_tensor.py | 5 +- src/eva/language/utils/text/messages.py | 52 +++- src/eva/multimodal/callbacks/__init__.py | 5 + .../multimodal/callbacks/writers/__init__.py | 5 + .../callbacks/writers/prediction.py | 39 +++ src/eva/multimodal/utils/text/messages.py | 6 +- .../assets/language/predictions/pubmedqa.csv | 3 + .../language/predictions/pubmedqa.jsonl | 3 + .../language/predictions/pubmedqa.parquet | 3 + .../callbacks/writers/test_prediction.py | 239 ++++++++++++++++++ .../language/data/datasets/test_prediction.py | 174 +++++++++++++ tests/eva/language/test_language_cli.py | 36 +++ .../models/modules/test_vision_language.py | 2 +- tests/eva/multimodal/test_multimodal_cli.py | 48 ++++ 39 files changed, 1304 insertions(+), 96 deletions(-) create mode 100644 configs/language/pathology/offline/multiple_choice/pubmedqa.yaml create mode 100644 configs/multimodal/pathology/offline/multiple_choice/patch_camelyon.yaml delete mode 100644 configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml create mode 100644 src/eva/language/callbacks/__init__.py create mode 100644 src/eva/language/callbacks/writers/__init__.py create mode 100644 src/eva/language/callbacks/writers/prediction.py create mode 100644 src/eva/language/data/datasets/prediction.py create mode 100644 src/eva/multimodal/callbacks/__init__.py create mode 100644 src/eva/multimodal/callbacks/writers/__init__.py create mode 100644 src/eva/multimodal/callbacks/writers/prediction.py create mode 100644 tests/eva/assets/language/predictions/pubmedqa.csv create mode 100644 tests/eva/assets/language/predictions/pubmedqa.jsonl create mode 100644 tests/eva/assets/language/predictions/pubmedqa.parquet create mode 100644 tests/eva/language/callbacks/writers/test_prediction.py create mode 100644 tests/eva/language/data/datasets/test_prediction.py diff --git a/.gitattributes b/.gitattributes index d4a50a98e..3f76ae8f2 100644 --- a/.gitattributes +++ b/.gitattributes @@ -3,6 +3,8 @@ tests/eva/assets/**/*.png filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.jpg filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.tif filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.tiff filter=lfs diff=lfs merge=lfs -text +tests/eva/assets/**/*.jsonl filter=lfs diff=lfs merge=lfs -text +tests/eva/assets/**/*.parquet filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.csv filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.pt filter=lfs diff=lfs merge=lfs -text tests/eva/assets/**/*.npy filter=lfs diff=lfs merge=lfs -text diff --git a/configs/language/pathology/offline/multiple_choice/pubmedqa.yaml b/configs/language/pathology/offline/multiple_choice/pubmedqa.yaml new file mode 100644 index 000000000..fb0aa18b0 --- /dev/null +++ b/configs/language/pathology/offline/multiple_choice/pubmedqa.yaml @@ -0,0 +1,74 @@ +--- +trainer: + class_path: eva.Trainer + init_args: + n_runs: &N_RUNS ${oc.env:N_RUNS, 1} + default_root_dir: ${oc.env:OUTPUT_ROOT, logs/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/language/pubmedqa} + checkpoint_type: null + callbacks: + - class_path: eva.callbacks.ConfigurationLogger + - class_path: eva.language.callbacks.writers.TextPredictionWriter + init_args: + output_dir: &PREDICTIONS_OUTPUT_DIR ${oc.env:PREDICTIONS_OUTPUT_DIR, ./predictions/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/pubmedqa} + dataloader_idx_map: + 0: val + save_format: &PREDICTIONS_SAVE_FORMAT ${oc.env:PREDICTIONS_SAVE_FORMAT, jsonl} + model: + class_path: eva.language.models.wrappers.ModelFromRegistry + init_args: + model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} + model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} + overwrite: false +model: + class_path: eva.language.models.OfflineLanguageModule + init_args: + metrics: + common: + - class_path: eva.metrics.MulticlassClassificationMetrics + init_args: + num_classes: 3 + input_type: "discrete" + postprocess: + predictions_transforms: + - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor + init_args: + mapping: {"no": 0, "yes": 1, "maybe": 2} + case_sensitive: false +data: + class_path: eva.DataModule + init_args: + datasets: + val: + class_path: eva.language.datasets.TextPredictionDataset + init_args: &DATASET_ARGS + path: ${oc.env:PREDICTIONS_OUTPUT_DIR, ./predictions/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/pubmedqa}/manifest.${oc.env:PREDICTIONS_SAVE_FORMAT, jsonl} + split: val + test: + class_path: eva.language.datasets.TextPredictionDataset + init_args: + <<: *DATASET_ARGS + predict: + - class_path: eva.language.datasets.PubMedQA + init_args: &PREDICT_DATASET_ARGS + root: ${oc.env:DATA_ROOT, ./data/pubmedqa} + split: val + download: ${oc.env:DOWNLOAD_DATA, false} + # Set `download: true` to download the dataset from https://huggingface.co/datasets/bigbio/pubmed_qa + # The PubMedQA dataset is distributed under the following license: MIT License + # See (https://github.com/pubmedqa/pubmedqa/blob/master/LICENSE) + - class_path: eva.language.datasets.PubMedQA + init_args: + <<: *PREDICT_DATASET_ARGS + split: test + dataloaders: + val: &DATALOADER_ARGS + batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 256} + num_workers: &N_DATA_WORKERS ${oc.env:N_DATA_WORKERS, 1} + shuffle: false + collate_fn: eva.language.data.dataloaders.prediction_collate + test: + <<: *DATALOADER_ARGS + predict: + batch_size: &PREDICT_BATCH_SIZE ${oc.env:PREDICT_BATCH_SIZE, 16} + num_workers: *N_DATA_WORKERS + collate_fn: eva.language.data.dataloaders.text_collate diff --git a/configs/language/pathology/online/multiple_choice/pubmedqa.yaml b/configs/language/pathology/online/multiple_choice/pubmedqa.yaml index cd322b82a..64eed22c5 100644 --- a/configs/language/pathology/online/multiple_choice/pubmedqa.yaml +++ b/configs/language/pathology/online/multiple_choice/pubmedqa.yaml @@ -5,6 +5,8 @@ trainer: n_runs: &N_RUNS ${oc.env:N_RUNS, 1} default_root_dir: ${oc.env:OUTPUT_ROOT, logs/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/pubmedqa} checkpoint_type: null + callbacks: + - class_path: eva.callbacks.ConfigurationLogger model: class_path: eva.language.models.LanguageModule init_args: @@ -33,15 +35,23 @@ data: class_path: eva.language.datasets.PubMedQA init_args: &DATASET_ARGS root: ${oc.env:DATA_ROOT, ./data/pubmedqa} - split: null + split: val download: ${oc.env:DOWNLOAD_DATA, false} # Set `download: true` to download the dataset from https://huggingface.co/datasets/bigbio/pubmed_qa # The PubMedQA dataset is distributed under the following license: MIT License # See (https://github.com/pubmedqa/pubmedqa/blob/master/LICENSE) - max_samples: 500 + test: + class_path: eva.language.datasets.PubMedQA + init_args: + <<: *DATASET_ARGS + split: test dataloaders: - val: + val: &DATALOADER_ARGS batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 16} num_workers: &N_DATA_WORKERS ${oc.env:N_DATA_WORKERS, 1} shuffle: false collate_fn: eva.language.data.dataloaders.text_collate + test: + <<: *DATALOADER_ARGS + + diff --git a/configs/multimodal/pathology/offline/multiple_choice/patch_camelyon.yaml b/configs/multimodal/pathology/offline/multiple_choice/patch_camelyon.yaml new file mode 100644 index 000000000..cd83f28c9 --- /dev/null +++ b/configs/multimodal/pathology/offline/multiple_choice/patch_camelyon.yaml @@ -0,0 +1,84 @@ +trainer: + class_path: eva.Trainer + init_args: + accelerator: ${oc.env:ACCELERATOR, auto} + n_runs: &N_RUNS ${oc.env:N_RUNS, 2} + default_root_dir: ${oc.env:OUTPUT_ROOT, logs/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/multimodal/patch_camelyon} + precision: bf16 + checkpoint_type: null + callbacks: + - class_path: eva.callbacks.ConfigurationLogger + - class_path: eva.multimodal.callbacks.writers.TextPredictionWriter + init_args: + output_dir: &PREDICTIONS_OUTPUT_DIR ${oc.env:PREDICTIONS_OUTPUT_DIR, ./predictions/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/multimodal/patch_camelyon} + dataloader_idx_map: + 0: val + 1: test + save_format: &PREDICTIONS_SAVE_FORMAT ${oc.env:PREDICTIONS_SAVE_FORMAT, jsonl} + model: + class_path: eva.multimodal.models.wrappers.ModelFromRegistry + init_args: + model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} + model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} + overwrite: false +model: + class_path: eva.language.models.OfflineLanguageModule + init_args: + metrics: + common: + - class_path: eva.metrics.MulticlassClassificationMetrics + init_args: + num_classes: 2 + input_type: "discrete" + postprocess: + predictions_transforms: + - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor + init_args: + mapping: {"A": 0, "B": 1} + case_sensitive: false +data: + class_path: eva.DataModule + init_args: + datasets: + val: + class_path: eva.language.datasets.TextPredictionDataset + init_args: &DATASET_ARGS + path: ${oc.env:PREDICTIONS_OUTPUT_DIR, ./predictions/${oc.env:MODEL_NAME, anthropic-claude-3-7-sonnet-latest}/multimodal/patch_camelyon}/manifest.${oc.env:PREDICTIONS_SAVE_FORMAT, jsonl} + split: val + test: + class_path: eva.language.datasets.TextPredictionDataset + init_args: + <<: *DATASET_ARGS + split: test + predict: + - class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: &PREDICTION_DATASET_ARGS + root: ${oc.env:DATA_ROOT, /mnt/localdisk/data/patch_camelyon} + split: val + download: ${oc.env:DOWNLOAD_DATA, false} + # Set `download: true` to download the dataset from https://zenodo.org/records/1494286 + # The PatchCamelyon dataset is distributed under the following license: + # "Creative Commons Zero v1.0 Universal" + # (see: https://choosealicense.com/licenses/cc0-1.0/) + transforms: + image: + class_path: eva.vision.data.transforms.Resize + init_args: + size: ${oc.env:RESIZE_DIM, null} + max_bytes: ${oc.env:IMAGE_MAX_BYTES, null} + max_samples: 500 + - class_path: eva.multimodal.data.datasets.PatchCamelyon + init_args: + <<: *PREDICTION_DATASET_ARGS + split: test + dataloaders: + val: &DATALOADER_ARGS + batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 256} + num_workers: &N_DATA_WORKERS ${oc.env:N_DATA_WORKERS, 1} + collate_fn: eva.language.data.dataloaders.prediction_collate + test: + <<: *DATALOADER_ARGS + predict: + batch_size: &PREDICT_BATCH_SIZE ${oc.env:PREDICT_BATCH_SIZE, 16} + num_workers: *N_DATA_WORKERS + collate_fn: eva.multimodal.data.dataloaders.text_image_collate diff --git a/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml b/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml deleted file mode 100644 index 15bdee4ee..000000000 --- a/configs/multimodal/tests/pathology/online/multiple_choice/patch_camelyon.yaml +++ /dev/null @@ -1,65 +0,0 @@ -trainer: - class_path: eva.Trainer - init_args: - accelerator: cpu - n_runs: &N_RUNS ${oc.env:N_RUNS, 2} - default_root_dir: &LIGHTNING_ROOT ${oc.env:LIGHTNING_ROOT, logs/test/multimodal/online/patch_camelyon} - max_epochs: &MAX_EPOCHS 1 - limit_val_batches: 2 - limit_test_batches: 2 - precision: bf16 - checkpoint_type: null - callbacks: - - class_path: eva.callbacks.ConfigurationLogger -model: - class_path: eva.multimodal.models.modules.VisionLanguageModule - init_args: - model: - class_path: eva.multimodal.models.wrappers.ModelFromRegistry - init_args: - model_name: ${oc.env:MODEL_NAME, anthropic/claude-3-7-sonnet-20250219} - model_extra_kwargs: ${oc.env:MODEL_EXTRA_KWARGS, null} - metrics: - common: - - class_path: eva.metrics.MulticlassClassificationMetrics - init_args: - num_classes: 2 - input_type: "discrete" - postprocess: - predictions_transforms: - - class_path: eva.language.utils.str_to_int_tensor.CastStrToIntTensor - init_args: - mapping: {"A": 0, "B": 1} - case_sensitive: false -data: - class_path: eva.DataModule - init_args: - datasets: - val: - class_path: eva.multimodal.data.datasets.PatchCamelyon - init_args: &DATASET_ARGS - root: ${oc.env:TESTS_ROOT, tests/eva}/assets/vision/datasets/patch_camelyon - split: val - download: false - transforms: - image: - class_path: eva.vision.data.transforms.Resize - init_args: - size: ${oc.env:RESIZE_DIM, null} - max_bytes: ${oc.env:IMAGE_MAX_BYTES, null} - max_samples: null - test: - class_path: eva.multimodal.data.datasets.PatchCamelyon - init_args: - <<: *DATASET_ARGS - split: test - dataloaders: - val: &DATALOADER_ARGS - batch_size: &BATCH_SIZE ${oc.env:BATCH_SIZE, 16} - collate_fn: eva.multimodal.data.dataloaders.text_image_collate - num_workers: 0 - pin_memory: false - persistent_workers: false - prefetch_factor: null - test: - <<: *DATALOADER_ARGS diff --git a/src/eva/core/interface/interface.py b/src/eva/core/interface/interface.py index 0ce8f7989..6091c8d2c 100644 --- a/src/eva/core/interface/interface.py +++ b/src/eva/core/interface/interface.py @@ -132,3 +132,24 @@ def test( n_runs=trainer.n_runs, verbose=trainer.n_runs > 1, ) + + def validate_test( + self, + trainer: eva_trainer.Trainer, + model: modules.ModelModule, + data: datamodules.DataModule, + ) -> None: + """Runs validation & test stages.""" + if getattr(data.datasets, "val", None) is None: + raise ValueError("The provided data module does not contain a validation dataset.") + if getattr(data.datasets, "test", None) is None: + raise ValueError("The provided data module does not contain a test dataset.") + + eva_trainer.run_evaluation_session( + base_trainer=trainer, + base_model=model, + datamodule=data, + stages=["validate", "test"], + n_runs=trainer.n_runs, + verbose=trainer.n_runs > 1, + ) diff --git a/src/eva/core/models/modules/module.py b/src/eva/core/models/modules/module.py index 55f3f2798..ac532fd59 100644 --- a/src/eva/core/models/modules/module.py +++ b/src/eva/core/models/modules/module.py @@ -33,8 +33,8 @@ def __init__( super().__init__() self._metrics = metrics or self.default_metrics - self._postprocess = postprocess or self.default_postprocess + self.postprocess = postprocess or self.default_postprocess self.metrics = metrics_lib.MetricModule.from_schema(self._metrics) @property @@ -133,7 +133,7 @@ def _common_batch_end(self, outputs: STEP_OUTPUT) -> STEP_OUTPUT: Returns: The updated outputs. """ - self._postprocess(outputs) + self.postprocess(outputs) return memory.recursive_detach(outputs, to_cpu=self.metrics_device.type == "cpu") def _forward_and_log_metrics( diff --git a/src/eva/language/callbacks/__init__.py b/src/eva/language/callbacks/__init__.py new file mode 100644 index 000000000..8d704af6a --- /dev/null +++ b/src/eva/language/callbacks/__init__.py @@ -0,0 +1,5 @@ +"""Language callbacks API.""" + +from eva.language.callbacks.writers import TextPredictionWriter + +__all__ = ["TextPredictionWriter"] diff --git a/src/eva/language/callbacks/writers/__init__.py b/src/eva/language/callbacks/writers/__init__.py new file mode 100644 index 000000000..9a5829cb8 --- /dev/null +++ b/src/eva/language/callbacks/writers/__init__.py @@ -0,0 +1,5 @@ +"""Language writers callbacks API.""" + +from eva.language.callbacks.writers.prediction import TextPredictionWriter + +__all__ = ["TextPredictionWriter"] diff --git a/src/eva/language/callbacks/writers/prediction.py b/src/eva/language/callbacks/writers/prediction.py new file mode 100644 index 000000000..9b29dd304 --- /dev/null +++ b/src/eva/language/callbacks/writers/prediction.py @@ -0,0 +1,176 @@ +"""Text prediction writer callbacks.""" + +import abc +import os +from typing import Any, Dict, List, Literal, Sequence, Tuple, TypedDict + +import lightning.pytorch as pl +import pandas as pd +import torch +from lightning.pytorch import callbacks +from torch import nn +from typing_extensions import NotRequired, override + +from eva.core.models.modules import utils as module_utils +from eva.language.models.typings import TextBatch +from eva.language.utils.text import messages as message_utils + + +class ManifestEntry(TypedDict): + """A single entry in the manifest file.""" + + prediction: str + """The predicted text.""" + + target: str + """The ground truth text.""" + + text: NotRequired[str] + """The input text data.""" + + split: NotRequired[str] + """The dataset split (e.g. train, val, test).""" + + +class TextPredictionWriter(callbacks.BasePredictionWriter, abc.ABC): + """Callback for writing generated text predictions to disk.""" + + def __init__( + self, + output_dir: str, + model: nn.Module, + dataloader_idx_map: Dict[int, str] | None = None, + metadata_keys: List[str] | None = None, + include_input: bool = True, + overwrite: bool = False, + apply_postprocess: bool = False, + save_format: Literal["jsonl", "parquet", "csv"] = "jsonl", + ) -> None: + """Initializes a new callback. + + Args: + output_dir: The directory where the embeddings will be saved. + model: The model instance used to generate the predictions. + dataloader_idx_map: A dictionary mapping dataloader indices to their respective + names (e.g. train, val, test). + metadata_keys: An optional list of keys to extract from the batch metadata and store + as additional columns in the manifest file. + include_input: Whether to include the original input text messages in the output. + overwrite: Whether to overwrite if embeddings are already present in the specified + output directory. If set to `False`, an error will be raised if embeddings are + already present (recommended). + apply_postprocess: Whether to apply the postprocesses specified in the model module. + save_format: The file format to use for saving the manifest file with the predictions. + """ + super().__init__() + self.output_dir = output_dir + self.model = model + self.dataloader_idx_map = dataloader_idx_map or {} + self.metadata_keys = metadata_keys + self.include_input = include_input + self.overwrite = overwrite + self.apply_postprocess = apply_postprocess + self.save_format = save_format + + self._manifest_path = os.path.join(self.output_dir, f"manifest.{self.save_format}") + self._data: List[ManifestEntry] = [] + + @override + def on_predict_start(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None: + self._check_if_exists() + + self.model = self.model.to(pl_module.device) + self.model.eval() + + @override + def write_on_batch_end( + self, + trainer: pl.Trainer, + pl_module: pl.LightningModule, + prediction: Any, + batch_indices: Sequence[int], + batch: TextBatch, + batch_idx: int, + dataloader_idx: int, + ) -> None: + text_batch, target_batch, metadata_batch = self._unpack_batch(batch) + has_target = target_batch is not None + split = self.dataloader_idx_map.get(dataloader_idx, "") + + prediction_batch = self._get_predictions(batch) + + target_batch, prediction_batch = self._apply_postprocess( + pl_module, target_batch, prediction_batch + ) + + for i in range(len(batch_indices)): + entry: ManifestEntry = { + "text": message_utils.serialize(text_batch[i]), + "prediction": str(prediction_batch[i]), + "target": str(target_batch[i]) if has_target else "", + "split": split if split else "", + } + + if self.metadata_keys is not None and metadata_batch is not None: + for key in self.metadata_keys: + entry[key] = metadata_batch[key][i] + + self._data.append(entry) + + @override + def on_predict_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None: + """Saves the gathered predictions to a manifest file.""" + df = pd.DataFrame(self._data) + + match self.save_format: + case "jsonl": + df.to_json(self._manifest_path, orient="records", lines=True) + case "parquet": + df.to_parquet(self._manifest_path, index=False) + case "csv": + df.to_csv(self._manifest_path, index=False) + case _: + raise ValueError(f"Unsupported save format: {self.save_format}") + + def _get_predictions(self, batch: TextBatch) -> List[str]: + with torch.no_grad(): + predictions = self.model(batch) + + if not isinstance(predictions, list) or not all(isinstance(p, str) for p in predictions): + raise ValueError("The model's output should be a list of strings.") + + return predictions + + def _check_if_exists(self) -> None: + """Checks if the output directory already exists and if it should be overwritten.""" + os.makedirs(self.output_dir, exist_ok=True) + if os.path.exists(self._manifest_path) and not self.overwrite: + raise FileExistsError( + f"The specified output directory already exists: {self.output_dir}. This " + "either means that the predictions have been computed before or that a " + "wrong output directory is being used." + ) + os.makedirs(self.output_dir, exist_ok=True) + + def _apply_postprocess( + self, pl_module: pl.LightningModule, targets: Any, predictions: Any + ) -> Tuple[List[Any], List[Any]]: + def _to_list(data: Any) -> List[Any]: + if isinstance(data, torch.Tensor): + return data.cpu().tolist() + return data + + if self.apply_postprocess and hasattr(pl_module, "postprocess"): + if ( + isinstance(pl_module.postprocess, module_utils.BatchPostProcess) + and pl_module.postprocess.predictions_transforms is not None + ): + outputs = {"targets": targets, "predictions": predictions} + pl_module.postprocess(outputs) + targets, predictions = outputs["targets"], outputs["predictions"] + + return _to_list(targets), _to_list(predictions) + + def _unpack_batch(self, batch: TextBatch) -> Tuple[list, list | None, dict | None]: + text_batch, target_batch, metadata_batch = TextBatch(*batch) + return text_batch, target_batch, metadata_batch diff --git a/src/eva/language/data/dataloaders/__init__.py b/src/eva/language/data/dataloaders/__init__.py index 90763136f..e5b2a045a 100644 --- a/src/eva/language/data/dataloaders/__init__.py +++ b/src/eva/language/data/dataloaders/__init__.py @@ -1,5 +1,5 @@ """Language Dataloaders API.""" -from eva.language.data.dataloaders.collate_fn import text_collate +from eva.language.data.dataloaders.collate_fn import prediction_collate, text_collate -__all__ = ["text_collate"] +__all__ = ["text_collate", "prediction_collate"] diff --git a/src/eva/language/data/dataloaders/collate_fn/__init__.py b/src/eva/language/data/dataloaders/collate_fn/__init__.py index 61671d160..7b471ffe7 100644 --- a/src/eva/language/data/dataloaders/collate_fn/__init__.py +++ b/src/eva/language/data/dataloaders/collate_fn/__init__.py @@ -1,5 +1,5 @@ """Collate functions API.""" -from eva.language.data.dataloaders.collate_fn.text import text_collate +from eva.language.data.dataloaders.collate_fn.text import prediction_collate, text_collate -__all__ = ["text_collate"] +__all__ = ["text_collate", "prediction_collate"] diff --git a/src/eva/language/data/dataloaders/collate_fn/text.py b/src/eva/language/data/dataloaders/collate_fn/text.py index d69a188e5..b2666447e 100644 --- a/src/eva/language/data/dataloaders/collate_fn/text.py +++ b/src/eva/language/data/dataloaders/collate_fn/text.py @@ -4,8 +4,8 @@ from torch.utils.data._utils.collate import default_collate -from eva.language.data.datasets.typings import TextSample -from eva.language.models.typings import TextBatch +from eva.language.data.datasets.typings import PredictionSample, TextSample +from eva.language.models.typings import PredictionBatch, TextBatch def text_collate(batch: List[TextSample]) -> TextBatch: @@ -30,3 +30,28 @@ def text_collate(batch: List[TextSample]) -> TextBatch: target=default_collate(targets) if targets[0] is not None else None, metadata=metadata, ) + + +def prediction_collate(batch: List[PredictionSample]) -> PredictionBatch: + """Collate function for text prediction data. + + Args: + batch: List of tuples containing (prediction, target, text, metadata) from the dataset + + Returns: + A batch of prediction samples. + """ + predictions, targets, texts, metadata = zip(*batch, strict=False) + first_sample = batch[0] + metadata = None + if first_sample.metadata is not None: + metadata = { + k: [sample.metadata[k] for sample in batch if sample.metadata] + for k in first_sample.metadata.keys() + } + return PredictionBatch( + prediction=default_collate(predictions) if predictions[0] is not None else None, + target=default_collate(targets) if targets[0] is not None else None, + text=list(texts) if first_sample.text is not None else None, + metadata=metadata, + ) diff --git a/src/eva/language/data/datasets/__init__.py b/src/eva/language/data/datasets/__init__.py index e51806933..3dd9fb8f8 100644 --- a/src/eva/language/data/datasets/__init__.py +++ b/src/eva/language/data/datasets/__init__.py @@ -2,8 +2,10 @@ from eva.language.data.datasets.base import LanguageDataset from eva.language.data.datasets.classification import PubMedQA +from eva.language.data.datasets.prediction import TextPredictionDataset __all__ = [ "PubMedQA", "LanguageDataset", + "TextPredictionDataset", ] diff --git a/src/eva/language/data/datasets/classification/pubmedqa.py b/src/eva/language/data/datasets/classification/pubmedqa.py index cd09bf88b..bf2b31681 100644 --- a/src/eva/language/data/datasets/classification/pubmedqa.py +++ b/src/eva/language/data/datasets/classification/pubmedqa.py @@ -16,6 +16,14 @@ class PubMedQA(base.TextClassification): """Dataset class for PubMedQA question answering task.""" + _expected_dataset_lengths: Dict[str | None, int] = { + "train": 450, + "val": 50, + "test": 500, + None: 500, + } + """Expected dataset lengths for the splits and complete dataset.""" + _license: str = "MIT License (https://github.com/pubmedqa/pubmedqa/blob/master/LICENSE)" """Dataset license.""" @@ -53,7 +61,14 @@ def _load_dataset(self, dataset_path: str | None) -> Dataset: """ dataset_name = "bigbio/pubmed_qa" config_name = "pubmed_qa_labeled_fold0_source" - split = (self._split or "train+test+validation") if self._split != "val" else "validation" + + match self._split: + case "val": + split = "validation" + case None: + split = "train+test+validation" + case _: + split = self._split if self._download: logger.info("Downloading dataset from HuggingFace Hub") @@ -89,7 +104,7 @@ def prepare_data(self) -> None: dataset_path = None if self._root: - dataset_path = self._root + dataset_path = os.path.join(self._root, self._split) if self._split else self._root os.makedirs(self._root, exist_ok=True) try: @@ -104,6 +119,15 @@ def prepare_data(self) -> None: except Exception as e: raise RuntimeError(f"Failed to prepare dataset: {e}") from e + @override + def validate(self) -> None: + if len(self) != self._expected_dataset_lengths[self._split]: + raise ValueError( + f"Dataset length mismatch for split '{self._split}': " + f"expected {self._expected_dataset_lengths[self._split]}, " + f"but got {len(self)}" + ) + @property @override def classes(self) -> List[str]: diff --git a/src/eva/language/data/datasets/prediction.py b/src/eva/language/data/datasets/prediction.py new file mode 100644 index 000000000..f13f15271 --- /dev/null +++ b/src/eva/language/data/datasets/prediction.py @@ -0,0 +1,151 @@ +"""Dataset class for loading pre-generated text predictions.""" + +import abc +from pathlib import Path +from typing import Any, Dict, Generic, Literal + +import pandas as pd +from typing_extensions import override + +from eva.language.data.datasets.base import LanguageDataset +from eva.language.data.datasets.schemas import TransformsSchema +from eva.language.data.datasets.typings import PredictionSample, TargetType +from eva.language.data.messages import MessageSeries, UserMessage +from eva.language.utils.text import messages as message_utils + + +class TextPredictionDataset( + LanguageDataset[PredictionSample[TargetType]], abc.ABC, Generic[TargetType] +): + """Dataset class for loading pre-generated text predictions.""" + + def __init__( + self, + path: str, + prediction_column: str = "prediction", + target_column: str = "target", + text_column: str | None = None, + metadata_columns: list[str] | None = None, + split: Literal["train", "val", "test"] | None = None, + transforms: TransformsSchema | None = None, + ): + """Initialize the dataset. + + Args: + path: The path to the manifest file holding the predictions & targets. + prediction_column: The name of the prediction column. + target_column: The name of the label column. + text_column: The name of the column with the text inputs that were used + to generate the predictions. If the text column contains chat message + json format ([{"role": ..., "content": ...}]), it will be deserialized into + a list of Message objects. Otherwise, the content is interpreted as a + single user message. + metadata_columns: List of column names to include in metadata. + split: The dataset split to use (train, val, test). If not specified, + the entire dataset will be used. + transforms: The transforms to apply to the text and target when + loading the samples. + """ + super().__init__() + + self.path = path + self.prediction_column = prediction_column + self.target_column = target_column + self.text_column = text_column + self.metadata_columns = metadata_columns + self.split = split + self.transforms = transforms + + self._data: pd.DataFrame + + @override + def __len__(self) -> int: + return len(self._data) + + @override + def __getitem__(self, index: int) -> PredictionSample[TargetType]: + item = PredictionSample( + prediction=self.load_prediction(index), + target=self.load_target(index), + text=self.load_text(index), + metadata=self.load_metadata(index) or {}, + ) + return self._apply_transforms(item) + + @override + def configure(self) -> None: + extension = Path(self.path).suffix + + match extension: + case ".jsonl": + self._data = pd.read_json(self.path, lines=True) + case ".csv": + self._data = pd.read_csv(self.path) + case ".parquet": + self._data = pd.read_parquet(self.path) + case _: + raise ValueError(f"Unsupported file extension: {extension}") + + if self.split is not None: + self._data = self._data[self._data["split"] == self.split].reset_index(drop=True) # type: ignore + + @override + def validate(self) -> None: + if self.prediction_column not in self._data.columns: + raise ValueError(f"Label column '{self.prediction_column}' not found.") + if self.target_column not in self._data.columns: + raise ValueError(f"Label column '{self.target_column}' not found.") + if self.metadata_columns: + missing_columns = set(self.metadata_columns) - set(self._data.columns) + if missing_columns: + raise ValueError(f"Metadata columns {missing_columns} not found.") + + def load_prediction(self, index: int) -> TargetType: + """Returns the prediction for the given index.""" + return self._data.iloc[index][self.prediction_column] + + def load_target(self, index: int) -> TargetType: + """Returns the target for the given index.""" + return self._data.iloc[index][self.target_column] + + def load_text(self, index: int) -> MessageSeries | None: + """Returns the text for the given index.""" + if self.text_column is None: + return None + + text = self._data.iloc[index][self.text_column] + + try: + return message_utils.deserialize(self._data.iloc[index][self.text_column]) + except Exception: + return [UserMessage(content=text)] + + def load_metadata(self, index: int) -> Dict[str, Any] | None: + """Returns the metadata for the given index.""" + if self.metadata_columns is None: + return None + + row = self._data.iloc[index] + return {col: row[col] for col in self.metadata_columns} + + def _apply_transforms( + self, sample: PredictionSample[TargetType] + ) -> PredictionSample[TargetType]: + """Applies the dataset transforms to the prediction and target.""" + if self.transforms: + text = self.transforms.text(sample.text) if self.transforms.text else sample.text + prediction = ( + self.transforms.prediction(sample.prediction) + if self.transforms.prediction + else sample.prediction + ) + target = ( + self.transforms.target(sample.target) if self.transforms.target else sample.target + ) + return PredictionSample( + prediction=prediction, + target=target, + text=text, + metadata=sample.metadata, + ) + return sample diff --git a/src/eva/language/data/datasets/schemas.py b/src/eva/language/data/datasets/schemas.py index 02359614e..4ea005497 100644 --- a/src/eva/language/data/datasets/schemas.py +++ b/src/eva/language/data/datasets/schemas.py @@ -13,3 +13,6 @@ class TransformsSchema: target: Callable | None = None """Target transformation""" + + prediction: Callable | None = None + """Prediction transformation""" diff --git a/src/eva/language/data/datasets/text.py b/src/eva/language/data/datasets/text.py index f2e3a71af..7a0019e78 100644 --- a/src/eva/language/data/datasets/text.py +++ b/src/eva/language/data/datasets/text.py @@ -75,7 +75,6 @@ def _apply_transforms(self, sample: TextSample[TargetType]) -> TextSample[Target Args: sample: The text sample.. - target: The target label. Returns: The transformed sample. diff --git a/src/eva/language/data/datasets/typings.py b/src/eva/language/data/datasets/typings.py index 0dd5b1b1f..edeffcc62 100644 --- a/src/eva/language/data/datasets/typings.py +++ b/src/eva/language/data/datasets/typings.py @@ -21,3 +21,19 @@ class TextSample(NamedTuple, Generic[TargetType]): metadata: dict[str, Any] | None """Additional metadata.""" + + +class PredictionSample(NamedTuple, Generic[TargetType]): + """Text sample with target and metadata.""" + + prediction: TargetType + """Prediction data.""" + + target: TargetType + """Target data.""" + + text: MessageSeries | None + """Conversation messages that were used as input.""" + + metadata: dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/language/data/messages.py b/src/eva/language/data/messages.py index 007d37797..4468d3936 100644 --- a/src/eva/language/data/messages.py +++ b/src/eva/language/data/messages.py @@ -1,9 +1,18 @@ """Types and classes for conversation messages in a multimodal context.""" import dataclasses +import enum from typing import Any, Dict, List +class Role(str, enum.Enum): + """Roles for messages in a conversation.""" + + USER = "user" + ASSISTANT = "assistant" + SYSTEM = "system" + + @dataclasses.dataclass class Message: """Base class for a message in a conversation.""" @@ -20,21 +29,21 @@ def to_dict(self) -> Dict[str, Any]: class UserMessage(Message): """User message in a conversation.""" - role: str = "user" + role: str = Role.USER @dataclasses.dataclass class AssistantMessage(Message): """Assistant message in a conversation.""" - role: str = "assistant" + role: str = Role.ASSISTANT @dataclasses.dataclass class SystemMessage(Message): """System message in a conversation.""" - role: str = "system" + role: str = Role.SYSTEM @dataclasses.dataclass diff --git a/src/eva/language/models/__init__.py b/src/eva/language/models/__init__.py index c8d3bf192..25ae94a40 100644 --- a/src/eva/language/models/__init__.py +++ b/src/eva/language/models/__init__.py @@ -1,7 +1,7 @@ """Language Models API.""" from eva.language.models import modules, networks, wrappers -from eva.language.models.modules import LanguageModule +from eva.language.models.modules import LanguageModule, OfflineLanguageModule from eva.language.models.wrappers import HuggingFaceModel, LiteLLMModel try: @@ -15,6 +15,7 @@ "LiteLLMModel", "VllmModel", "LanguageModule", + "OfflineLanguageModule", ] except ImportError: __all__ = [ @@ -24,4 +25,5 @@ "HuggingFaceModel", "LiteLLMModel", "LanguageModule", + "OfflineLanguageModule", ] diff --git a/src/eva/language/models/modules/__init__.py b/src/eva/language/models/modules/__init__.py index 3dcd2cfc2..cbf78541d 100644 --- a/src/eva/language/models/modules/__init__.py +++ b/src/eva/language/models/modules/__init__.py @@ -1,5 +1,5 @@ """Language Networks API.""" -from eva.language.models.modules.language import LanguageModule +from eva.language.models.modules.language import LanguageModule, OfflineLanguageModule -__all__ = ["LanguageModule"] +__all__ = ["LanguageModule", "OfflineLanguageModule"] diff --git a/src/eva/language/models/modules/language.py b/src/eva/language/models/modules/language.py index 74cb3ef87..f2f4bae47 100644 --- a/src/eva/language/models/modules/language.py +++ b/src/eva/language/models/modules/language.py @@ -9,7 +9,7 @@ from eva.core.metrics import structs as metrics_lib from eva.core.models.modules import module from eva.core.models.modules.utils import batch_postprocess -from eva.language.models.typings import TextBatch +from eva.language.models.typings import PredictionBatch, TextBatch class LanguageModule(module.ModelModule): @@ -53,3 +53,41 @@ def _batch_step(self, batch: TextBatch) -> STEP_OUTPUT: "targets": targets, "metadata": metadata, } + + +class OfflineLanguageModule(module.ModelModule): + """Model module for offline language tasks.""" + + def __init__( + self, + metrics: metrics_lib.MetricsSchema | None = None, + postprocess: batch_postprocess.BatchPostProcess | None = None, + ) -> None: + """Initializes the text inference module. + + Args: + metrics: Metrics schema for evaluation. + postprocess: A helper function to post-process model outputs before evaluation. + """ + super().__init__(metrics=metrics, postprocess=postprocess) + + @override + def forward(self, batch: PredictionBatch, *args: Any, **kwargs: Any) -> PredictionBatch: + return batch + + @override + def validation_step(self, batch: PredictionBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + @override + def test_step(self, batch: PredictionBatch, *args: Any, **kwargs: Any) -> STEP_OUTPUT: + return self._batch_step(batch) + + def _batch_step(self, batch: PredictionBatch) -> STEP_OUTPUT: + predictions, targets, text, metadata = PredictionBatch(*batch) + return { + "inputs": text, + "predictions": predictions, + "targets": targets, + "metadata": metadata, + } diff --git a/src/eva/language/models/typings.py b/src/eva/language/models/typings.py index 71b35f019..74db85019 100644 --- a/src/eva/language/models/typings.py +++ b/src/eva/language/models/typings.py @@ -21,3 +21,19 @@ class TextBatch(NamedTuple, Generic[TargetType]): metadata: Dict[str, Any] | None """Additional metadata.""" + + +class PredictionBatch(NamedTuple, Generic[TargetType]): + """Text sample with target and metadata.""" + + prediction: TargetType + """Prediction data.""" + + target: TargetType + """Target data.""" + + text: List[MessageSeries] | None + """Conversation messages that were used as input.""" + + metadata: Dict[str, Any] | None + """Additional metadata.""" diff --git a/src/eva/language/models/wrappers/__init__.py b/src/eva/language/models/wrappers/__init__.py index 00482d826..53a1b2add 100644 --- a/src/eva/language/models/wrappers/__init__.py +++ b/src/eva/language/models/wrappers/__init__.py @@ -1,5 +1,6 @@ """Language Model Wrappers API.""" +from eva.language.models.wrappers.base import LanguageModel from eva.language.models.wrappers.from_registry import ModelFromRegistry from eva.language.models.wrappers.huggingface import HuggingFaceModel from eva.language.models.wrappers.litellm import LiteLLMModel @@ -7,6 +8,12 @@ try: from eva.language.models.wrappers.vllm import VllmModel - __all__ = ["HuggingFaceModel", "LiteLLMModel", "VllmModel", "ModelFromRegistry"] + __all__ = [ + "LanguageModel", + "HuggingFaceModel", + "LiteLLMModel", + "VllmModel", + "ModelFromRegistry", + ] except ImportError: - __all__ = ["HuggingFaceModel", "LiteLLMModel", "ModelFromRegistry"] + __all__ = ["LanguageModel", "HuggingFaceModel", "LiteLLMModel", "ModelFromRegistry"] diff --git a/src/eva/language/utils/str_to_int_tensor.py b/src/eva/language/utils/str_to_int_tensor.py index 67e2977cd..120c62f9a 100644 --- a/src/eva/language/utils/str_to_int_tensor.py +++ b/src/eva/language/utils/str_to_int_tensor.py @@ -63,7 +63,10 @@ def __call__(self, values: Union[str, List[str], List[int]]) -> torch.Tensor: ValueError: If any value cannot be mapped to an integer. """ return torch.tensor( - [self._cast_single(v) for v in (values if isinstance(values, list) else [values])], + [ + self._cast_single(v) + for v in (values if isinstance(values, list | tuple) else [values]) + ], dtype=torch.int, ) diff --git a/src/eva/language/utils/text/messages.py b/src/eva/language/utils/text/messages.py index 89ba753bb..8b57bd1bc 100644 --- a/src/eva/language/utils/text/messages.py +++ b/src/eva/language/utils/text/messages.py @@ -1,9 +1,16 @@ """Message formatting utilities for language models.""" import functools +import json from typing import Any, Dict, List -from eva.language.data.messages import MessageSeries, SystemMessage +from eva.language.data.messages import ( + AssistantMessage, + MessageSeries, + Role, + SystemMessage, + UserMessage, +) def format_chat_message(message: MessageSeries) -> List[Dict[str, Any]]: @@ -26,11 +33,11 @@ def combine_system_messages(message: MessageSeries, join_char: str = "\n") -> Me A new message series with system messages combined into one and the remaining messages unchanged. """ - system_messages = list(filter(lambda item: item.role == "system", message)) + system_messages = list(filter(lambda item: item.role == Role.SYSTEM, message)) if len(system_messages) == 0: return message - non_system_messages = list(filter(lambda item: item.role != "system", message)) + non_system_messages = list(filter(lambda item: item.role != Role.SYSTEM, message)) return [ SystemMessage(content=merge_message_contents(system_messages, join_char=join_char)) ] + non_system_messages @@ -65,3 +72,42 @@ def batch_insert_system_message( return list( map(functools.partial(insert_system_message, system_message=system_message), messages) ) + + +def serialize(messages: MessageSeries) -> str: + """Serialize a MessageSeries object into a JSON string. + + Args: + messages: A list of message objects (MessagesSeries). + + Returns: + A JSON string representing the message series, with the following format: + [{"role": "user", "content": "Hello"}, ...] + """ + serialized_messages = format_chat_message(messages) + return json.dumps(serialized_messages) + + +def deserialize(messages: str) -> MessageSeries: + """Convert a json string to a MessageSeries object. + + Format: [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "Hi there!"}] + """ + message_dicts = json.loads(messages) + + message_series = [] + for message_dict in message_dicts: + if "role" not in message_dict or "content" not in message_dict: + raise ValueError("`role` or `content` keys are missing.") + + match message_dict["role"]: + case Role.USER: + message_series.append(UserMessage(**message_dict)) + case Role.ASSISTANT: + message_series.append(AssistantMessage(**message_dict)) + case Role.SYSTEM: + message_series.append(SystemMessage(**message_dict)) + case _: + raise ValueError(f"Unknown role: {message_dict['role']}") + + return message_series diff --git a/src/eva/multimodal/callbacks/__init__.py b/src/eva/multimodal/callbacks/__init__.py new file mode 100644 index 000000000..4a0a614b8 --- /dev/null +++ b/src/eva/multimodal/callbacks/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal callbacks API.""" + +from eva.multimodal.callbacks.writers import TextPredictionWriter + +__all__ = ["TextPredictionWriter"] diff --git a/src/eva/multimodal/callbacks/writers/__init__.py b/src/eva/multimodal/callbacks/writers/__init__.py new file mode 100644 index 000000000..e12327ee7 --- /dev/null +++ b/src/eva/multimodal/callbacks/writers/__init__.py @@ -0,0 +1,5 @@ +"""Multimodal writers callbacks API.""" + +from eva.multimodal.callbacks.writers.prediction import TextPredictionWriter + +__all__ = ["TextPredictionWriter"] diff --git a/src/eva/multimodal/callbacks/writers/prediction.py b/src/eva/multimodal/callbacks/writers/prediction.py new file mode 100644 index 000000000..c074d7fb8 --- /dev/null +++ b/src/eva/multimodal/callbacks/writers/prediction.py @@ -0,0 +1,39 @@ +"""Text prediction writer callbacks.""" + +from typing import Dict, List, Literal, Tuple + +from torch import nn +from typing_extensions import override + +from eva.language.callbacks import writers +from eva.multimodal.models.typings import TextImageBatch + + +class TextPredictionWriter(writers.TextPredictionWriter): + """Callback for writing generated text predictions to disk.""" + + def __init__( + self, + output_dir: str, + model: nn.Module, + dataloader_idx_map: Dict[int, str] | None = None, + metadata_keys: List[str] | None = None, + include_input: bool = True, + overwrite: bool = False, + save_format: Literal["jsonl", "parquet", "csv"] = "jsonl", + ) -> None: + """See docstring of base class.""" + super().__init__( + output_dir=output_dir, + model=model, + dataloader_idx_map=dataloader_idx_map, + metadata_keys=metadata_keys, + include_input=include_input, + overwrite=overwrite, + save_format=save_format, + ) + + @override + def _unpack_batch(self, batch: TextImageBatch) -> Tuple[list, list | None, dict | None]: # type: ignore + text_batch, _, target_batch, metadata_batch = TextImageBatch(*batch) + return text_batch, target_batch, metadata_batch diff --git a/src/eva/multimodal/utils/text/messages.py b/src/eva/multimodal/utils/text/messages.py index 573e9cad1..457a14212 100644 --- a/src/eva/multimodal/utils/text/messages.py +++ b/src/eva/multimodal/utils/text/messages.py @@ -5,7 +5,7 @@ from torchvision import tv_tensors from eva.language import utils as language_utils -from eva.language.data.messages import MessageSeries +from eva.language.data.messages import MessageSeries, Role from eva.multimodal.utils import image as image_utils @@ -18,7 +18,7 @@ def format_huggingface_message( formatted_message = [] for item in message: - if item.role == "system": + if item.role == Role.SYSTEM: formatted_message += language_utils.format_chat_message([item]) else: formatted_message.append( @@ -53,7 +53,7 @@ def format_litellm_message( formatted_message = [] for item in message: - if item.role == "system": + if item.role == Role.SYSTEM: formatted_message += language_utils.format_chat_message([item]) else: formatted_message.append( diff --git a/tests/eva/assets/language/predictions/pubmedqa.csv b/tests/eva/assets/language/predictions/pubmedqa.csv new file mode 100644 index 000000000..c8370ba38 --- /dev/null +++ b/tests/eva/assets/language/predictions/pubmedqa.csv @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:39744f2aa8c9e25507d60f6506a40be025bdb6e256467194d6a6cf0e103ff6e0 +size 13550 diff --git a/tests/eva/assets/language/predictions/pubmedqa.jsonl b/tests/eva/assets/language/predictions/pubmedqa.jsonl new file mode 100644 index 000000000..ab1f5b78d --- /dev/null +++ b/tests/eva/assets/language/predictions/pubmedqa.jsonl @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:014c4f2a91880599a16486fdbcee0323955c34330460010ed37d5ae294b9efbe +size 13907 diff --git a/tests/eva/assets/language/predictions/pubmedqa.parquet b/tests/eva/assets/language/predictions/pubmedqa.parquet new file mode 100644 index 000000000..8d67da466 --- /dev/null +++ b/tests/eva/assets/language/predictions/pubmedqa.parquet @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:679005ced11905c0302654f162ae7c70447124f4151d16cdca3e8cbabae8c483 +size 13505 diff --git a/tests/eva/language/callbacks/writers/test_prediction.py b/tests/eva/language/callbacks/writers/test_prediction.py new file mode 100644 index 000000000..2f76c0537 --- /dev/null +++ b/tests/eva/language/callbacks/writers/test_prediction.py @@ -0,0 +1,239 @@ +"""Tests the text prediction writer.""" + +import functools +import os +import tempfile +from typing import List, Literal, cast + +import lightning.pytorch as pl +import pandas as pd +import pytest +from lightning.pytorch import Callback +from lightning.pytorch.demos import boring_classes +from torch import nn +from typing_extensions import override + +from eva.core.data import dataloaders, datamodules, datasets +from eva.language.callbacks.writers import prediction as prediction_writer +from eva.language.data.dataloaders import text_collate +from eva.language.data.datasets.typings import TextSample +from eva.language.data.messages import UserMessage + + +@pytest.mark.parametrize( + "batch_size, n_samples, save_format, include_input", + [ + (5, 7, "jsonl", True), + (8, 16, "csv", True), + (8, 32, "parquet", True), + (5, 7, "jsonl", False), + ], +) +def test_prediction_writer( + datamodule: datamodules.DataModule, + model: pl.LightningModule, + text_generation_model: nn.Module, + save_format: Literal["jsonl", "parquet", "csv"], + include_input: bool, +) -> None: + """Tests the text prediction writer callback. + + This test executes a lightning trainer predict operation and checks if the expected + predictions are correctly written to disk. + """ + with tempfile.TemporaryDirectory() as output_dir: + callback = prediction_writer.TextPredictionWriter( + output_dir=output_dir, + model=text_generation_model, + dataloader_idx_map={0: "train", 1: "val", 2: "test"}, + metadata_keys=["example_metadata"], + include_input=include_input, + overwrite=True, + save_format=save_format, + ) + trainer = _init_and_run_trainer([callback], model, datamodule) + + assert isinstance(trainer.predict_dataloaders, list) + assert len(trainer.predict_dataloaders) == 3 + + _check_manifest(output_dir, datamodule, save_format, include_input) + + +@pytest.mark.parametrize("batch_size, n_samples, save_format", [(5, 7, "jsonl")]) +def test_prediction_writer_overwrite_protection( + datamodule: datamodules.DataModule, + model: pl.LightningModule, + text_generation_model: nn.Module, + save_format: Literal["jsonl", "parquet", "csv"], +) -> None: + """Tests that the writer raises an error when overwrite is False and files exist.""" + with tempfile.TemporaryDirectory() as output_dir: + callback = prediction_writer.TextPredictionWriter( + output_dir=output_dir, + model=text_generation_model, + overwrite=True, + save_format=save_format, + ) + _init_and_run_trainer([callback], model, datamodule) + + # Try to write again without overwrite + callback2 = prediction_writer.TextPredictionWriter( + output_dir=output_dir, + model=text_generation_model, + overwrite=False, + save_format=save_format, + ) + + with pytest.raises(FileExistsError): + _init_and_run_trainer([callback2], model, datamodule) + + +def _init_and_run_trainer( + callbacks: List[Callback], + model: pl.LightningModule, + datamodule: datamodules.DataModule, +): + """Initializes and runs the trainer with the given callbacks.""" + trainer = pl.Trainer( + logger=False, + accelerator="cpu", + callbacks=callbacks, + ) + trainer.predict(model=model, datamodule=datamodule, return_predictions=True) + + return trainer + + +def _check_manifest( + output_dir: str, + datamodule: datamodules.DataModule, + save_format: Literal["jsonl", "parquet", "csv"], + include_input: bool = True, +): + """Checks if the manifest file contains the expected entries.""" + manifest_path = os.path.join(output_dir, f"manifest.{save_format}") + assert os.path.isfile(manifest_path) + + # Load manifest based on format + match save_format: + case "jsonl": + df_manifest = pd.read_json(manifest_path, lines=True) + case "csv": + df_manifest = pd.read_csv(manifest_path) + case "parquet": + df_manifest = pd.read_parquet(manifest_path) + case _: + raise ValueError(f"Unsupported save format: {save_format}") + + # Check expected columns + expected_columns = ["text", "prediction", "target", "split", "example_metadata"] + for column in expected_columns: + assert column in df_manifest.columns + + # Check number of entries + total_samples = sum(len(ds) for ds in datamodule.datasets.predict) # type: ignore + assert len(df_manifest) == total_samples + + # Check that predictions are strings + assert all(isinstance(pred, str) for pred in df_manifest["prediction"]) + + # Check that splits are correctly assigned + assert set(df_manifest["split"]) == {"train", "val", "test"} + + +@pytest.fixture(scope="function") +def text_generation_model() -> nn.Module: + """Returns a simple model that generates text predictions.""" + return FakeTextModel() + + +@pytest.fixture(scope="function") +def model(text_generation_model: nn.Module) -> pl.LightningModule: + """Returns a LightningModule wrapper for the text model.""" + return FakeLightningModule(text_generation_model) + + +@pytest.fixture(scope="function") +def dataset(n_samples: int) -> List[datasets.TorchDataset]: + """Fake dataset fixture.""" + Dataset = functools.partial( + FakeTextDataset, + length=n_samples, + ) + train_dataset = Dataset(split="train") + val_dataset = Dataset(split="val") + test_dataset = Dataset(split="test") + + return [train_dataset, val_dataset, test_dataset] + + +@pytest.fixture(scope="function") +def datamodule(batch_size: int, dataset: List[datasets.TorchDataset]) -> datamodules.DataModule: + """Returns a DataModule fixture.""" + dataloader = dataloaders.DataLoader( + batch_size=batch_size, + num_workers=0, + pin_memory=False, + persistent_workers=False, + prefetch_factor=None, + collate_fn=text_collate, + ) + return datamodules.DataModule( + datasets=datamodules.DatasetsSchema( + train=dataset[0], val=dataset[1], predict=cast(List[datasets.TorchDataset], dataset) + ), + dataloaders=datamodules.DataloadersSchema( + train=dataloader, + val=dataloader, + predict=dataloader, + ), + ) + + +class FakeTextModel(nn.Module): + """Fake model that generates text predictions.""" + + def forward(self, batch): + """Returns a list of fake predictions.""" + text_batch, _, _ = batch + batch_size = len(text_batch) + predictions = [f"prediction_{i}" for i in range(batch_size)] + return predictions + + +class FakeLightningModule(pl.LightningModule): + """Fake LightningModule wrapper.""" + + def __init__(self, text_model: nn.Module): + """Initializes the module.""" + super().__init__() + self.text_model = text_model + + def forward(self, x): + """Forward pass.""" + return self.text_model(x) + + def predict_step(self, batch, batch_idx, dataloader_idx=0): + """Prediction step.""" + return self(batch) + + +class FakeTextDataset(boring_classes.RandomDataset, datasets.Dataset): + """Fake text dataset.""" + + def __init__( + self, + split: Literal["train", "val", "test"], + length: int = 10, + ): + """Initializes the dataset.""" + super().__init__(size=32, length=length) + self._split = split + + @override + def __getitem__(self, index: int): + """Returns a text sample with metadata.""" + text = [UserMessage(content=f"Sample text {self._split}-{index}")] + target = f"target_{index}" + metadata = {"example_metadata": f"metadata_{index}"} + return TextSample(text=text, target=target, metadata=metadata) # type: ignore diff --git a/tests/eva/language/data/datasets/test_prediction.py b/tests/eva/language/data/datasets/test_prediction.py new file mode 100644 index 000000000..b4945768c --- /dev/null +++ b/tests/eva/language/data/datasets/test_prediction.py @@ -0,0 +1,174 @@ +"""TextPredictionDataset tests.""" + +from pathlib import Path + +import pytest + +from eva.language.data.datasets.prediction import TextPredictionDataset +from eva.language.data.datasets.schemas import TransformsSchema +from eva.language.data.datasets.typings import PredictionSample +from eva.language.data.messages import UserMessage + + +@pytest.mark.parametrize( + "file_format, expected_length", + [("jsonl", 8), ("csv", 8), ("parquet", 8)], +) +def test_length(prediction_dataset: TextPredictionDataset, expected_length: int) -> None: + """Tests the length of the dataset.""" + assert len(prediction_dataset) == expected_length + + +@pytest.mark.parametrize( + "file_format, index", + [ + ("jsonl", 0), + ("jsonl", 5), + ("csv", 0), + ("parquet", 0), + ], +) +def test_sample(prediction_dataset: TextPredictionDataset, index: int) -> None: + """Tests the format of a dataset sample.""" + sample = prediction_dataset[index] + + assert isinstance(sample, PredictionSample) + assert sample.prediction in ["yes", "no", "maybe"] + # Target can be string or int depending on the file format + assert sample.target in ["0", "1", "2", 0, 1, 2] + + assert isinstance(sample.text, list) + assert len(sample.text) == 1 + assert isinstance(sample.text[0], UserMessage) + assert "Question:" in sample.text[0].content + + assert isinstance(sample.metadata, dict) + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_no_text_column(assets_path: Path, file_format: str) -> None: + """Tests dataset without text column.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="target", + text_column=None, + ) + dataset.setup() + + sample = dataset[0] + assert sample.text is None + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_with_split(assets_path: Path, file_format: str) -> None: + """Tests dataset with split filtering.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="target", + text_column="text", + split="val", + ) + dataset.setup() + + assert len(dataset) == 4 + sample = dataset[0] + assert isinstance(sample, PredictionSample) + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_with_metadata_columns(assets_path: Path, file_format: str) -> None: + """Tests dataset with metadata columns.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="target", + text_column="text", + metadata_columns=["split"], + ) + dataset.setup() + + sample = dataset[0] + assert isinstance(sample.metadata, dict) + assert "split" in sample.metadata + assert sample.metadata["split"] == "val" + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_unsupported_format(assets_path: Path, file_format: str) -> None: + """Tests configuration with unsupported file format.""" + dataset = TextPredictionDataset( + path="dummy.txt", + prediction_column="prediction", + target_column="target", + ) + with pytest.raises(ValueError, match="Unsupported file extension"): + dataset.configure() + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_missing_columns(assets_path: Path, file_format: str) -> None: + """Tests validation with missing columns.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + + # Missing prediction column + dataset = TextPredictionDataset( + path=str(path), + prediction_column="non_existent", + target_column="target", + ) + dataset.configure() + with pytest.raises(ValueError, match="Label column 'non_existent' not found"): + dataset.validate() + + # Missing target column + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="non_existent", + ) + dataset.configure() + with pytest.raises(ValueError, match="Label column 'non_existent' not found"): + dataset.validate() + + +@pytest.mark.parametrize("file_format", ["jsonl"]) +def test_with_transforms(assets_path: Path, file_format: str) -> None: + """Tests dataset with transforms.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + + transforms = TransformsSchema( + prediction=lambda x: f"transformed_{x}", + target=lambda x: int(x) + 100, + ) + + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="target", + text_column="text", + transforms=transforms, + ) + dataset.setup() + + sample = dataset[0] + assert sample.prediction.startswith("transformed_") + assert sample.target == 101 + + +@pytest.fixture(scope="function") +def prediction_dataset(assets_path: Path, file_format: str) -> TextPredictionDataset: + """TextPredictionDataset fixture.""" + path = Path(assets_path) / "language" / "predictions" / f"pubmedqa.{file_format}" + dataset = TextPredictionDataset( + path=str(path), + prediction_column="prediction", + target_column="target", + text_column="text", + ) + dataset.setup() + return dataset diff --git a/tests/eva/language/test_language_cli.py b/tests/eva/language/test_language_cli.py index d8ceb68f5..421bb2bc6 100644 --- a/tests/eva/language/test_language_cli.py +++ b/tests/eva/language/test_language_cli.py @@ -1,6 +1,7 @@ """Tests regarding eva's CLI commands on language datasets.""" import os +import tempfile from unittest import mock from unittest.mock import patch @@ -15,6 +16,7 @@ "configuration_file", [ "configs/language/pathology/online/multiple_choice/pubmedqa.yaml", + "configs/language/pathology/offline/multiple_choice/pubmedqa.yaml", ], ) def test_configuration_initialization(configuration_file: str, lib_path: str) -> None: @@ -47,6 +49,34 @@ def test_validate_from_configuration(configuration_file: str, lib_path: str) -> ) +@pytest.mark.parametrize( + "configuration_file", + [ + "configs/language/pathology/offline/multiple_choice/pubmedqa.yaml", + ], +) +def test_predict_validate_from_configuration(configuration_file: str, lib_path: str) -> None: + """Tests CLI `predict` and `validate` commands with a given configuration file.""" + with tempfile.TemporaryDirectory() as output_dir: + with mock.patch.dict( + os.environ, {"N_RUNS": "1", "BATCH_SIZE": "2", "PREDICTIONS_OUTPUT_DIR": output_dir} + ): + _cli.run_cli_from_main( + cli_args=[ + "predict", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + _cli.run_cli_from_main( + cli_args=[ + "validate", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + + @pytest.fixture(autouse=True) def mock_dependencies(): """Mocks external dependencies to avoid API calls and downloads.""" @@ -74,3 +104,9 @@ def _fake_prepare_data(self): mock.patch.dict(os.environ, {"ANTHROPIC_API_KEY": "dummy-key"}), ): yield + + +@pytest.fixture(autouse=True) +def skip_dataset_validation() -> None: + """Mocks the validation step of the datasets.""" + datasets.PubMedQA.validate = mock.MagicMock(return_value=None) diff --git a/tests/eva/multimodal/models/modules/test_vision_language.py b/tests/eva/multimodal/models/modules/test_vision_language.py index 51f9fd251..5c0bdb715 100644 --- a/tests/eva/multimodal/models/modules/test_vision_language.py +++ b/tests/eva/multimodal/models/modules/test_vision_language.py @@ -75,7 +75,7 @@ def test_init_attributes(model): module_instance = VisionLanguageModule(model=model) assert module_instance.model is model assert module_instance.metrics is not None # MetricModule is created by default - assert module_instance._postprocess is not None # BatchPostProcess is created by default + assert module_instance.postprocess is not None # BatchPostProcess is created by default def test_batch_step_without_targets(vision_language_module): diff --git a/tests/eva/multimodal/test_multimodal_cli.py b/tests/eva/multimodal/test_multimodal_cli.py index a63a80592..b10e27abb 100644 --- a/tests/eva/multimodal/test_multimodal_cli.py +++ b/tests/eva/multimodal/test_multimodal_cli.py @@ -1,6 +1,7 @@ """Tests regarding eva's CLI commands on multimodal datasets.""" import os +import tempfile from unittest import mock from unittest.mock import patch @@ -48,6 +49,53 @@ def test_validate_from_configuration(configuration_file: str, lib_path: str) -> ) +@pytest.mark.parametrize( + "configuration_file", + [ + "configs/multimodal/pathology/online/multiple_choice/patch_camelyon.yaml", + ], +) +def test_test_from_configuration(configuration_file: str, lib_path: str) -> None: + """Tests CLI `test` command with a given configuration file.""" + with mock.patch.dict(os.environ, {"N_RUNS": "1", "BATCH_SIZE": f"{BATCH_SIZE}"}): + _cli.run_cli_from_main( + cli_args=[ + "test", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + + +@pytest.mark.parametrize( + "configuration_file", + [ + "configs/multimodal/pathology/offline/multiple_choice/patch_camelyon.yaml", + ], +) +def test_predict_validate_from_configuration(configuration_file: str, lib_path: str) -> None: + """Tests CLI `predict` and `validate` commands with a given configuration file.""" + with tempfile.TemporaryDirectory() as output_dir: + with mock.patch.dict( + os.environ, + {"N_RUNS": "1", "PREDICT_BATCH_SIZE": "2", "PREDICTIONS_OUTPUT_DIR": output_dir}, + ): + _cli.run_cli_from_main( + cli_args=[ + "predict", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + _cli.run_cli_from_main( + cli_args=[ + "validate", + "--config", + os.path.join(lib_path, configuration_file), + ] + ) + + @pytest.fixture(autouse=True) def skip_dataset_validation() -> None: """Mocks the validation step of the datasets."""