From 02501d5d199c721ade5caf4fee0f0a7396b98d61 Mon Sep 17 00:00:00 2001 From: MaartenGr Date: Thu, 27 Aug 2026 18:49:48 +0200 Subject: [PATCH] Support multimodal in Corpus --- bertopic/_bertopic.py | 32 ++++++++-- bertopic/_corpus.py | 88 +++++++++++++++++++------- tests/test_corpus.py | 141 +++++++++++++++++++++++++++++++++++++----- 3 files changed, 217 insertions(+), 44 deletions(-) diff --git a/bertopic/_bertopic.py b/bertopic/_bertopic.py index 80c56761..0fcacb15 100644 --- a/bertopic/_bertopic.py +++ b/bertopic/_bertopic.py @@ -66,7 +66,7 @@ from bertopic.vectorizers import ClassTfidfTransformer from bertopic.representation import BaseRepresentation from bertopic._topics import Keywords, Topic, Topics, TopicRepresentation, TopicHierarchy -from bertopic._corpus import Corpus +from bertopic._corpus import Corpus, Modality from bertopic.cluster._utils import hdbscan_delegator, is_supported_hdbscan from bertopic._utils import ( MyLogger, @@ -458,7 +458,13 @@ def fit_transform( topics, probs = topic_model.fit_transform(docs, embeddings) ``` """ - corpus = Corpus(documents=documents, embeddings=embeddings, images=images, y=y) + corpus = Corpus( + documents=documents, + media=images, + modality=Modality.IMAGE if images else Modality.TEXT, + embeddings=embeddings, + y=y, + ) # 1. Extract embeddings if corpus.embeddings is None or (corpus.embeddings is not None and self.embedding_model is not None): @@ -598,7 +604,12 @@ def transform( ``` """ check_is_fitted(self) - corpus = Corpus(documents=documents, embeddings=embeddings, images=images) + corpus = Corpus( + documents=documents, + media=images, + modality=Modality.IMAGE if images else Modality.TEXT, + embeddings=embeddings, + ) # Extract embeddings if corpus.embeddings is None: @@ -1051,7 +1062,13 @@ def update_topics( # (duplicate topic embedding for each document based on assignment) topic_embeddings = {topic.id: topic.embedding for topic in self._topics} doc_embeddings = np.array([topic_embeddings[topic_id] for topic_id in topics]) - corpus = Corpus(documents=docs, topics=np.array(topics), images=images, embeddings=doc_embeddings) + corpus = Corpus( + documents=docs, + media=images, + modality=Modality.IMAGE if images else Modality.TEXT, + topics=np.array(topics), + embeddings=doc_embeddings, + ) self._extract_representations(corpus, fine_tune=True) self._save_representative_docs(corpus) @@ -1571,8 +1588,9 @@ def merge_topics( doc_embeddings = np.array([topic_embeddings[topic_id] for topic_id in self.topics_]) corpus = Corpus( documents=docs, + media=images, + modality=Modality.IMAGE if images else Modality.TEXT, topics=np.array(self.topics_), - images=images, embeddings=doc_embeddings, ) @@ -1661,8 +1679,9 @@ def reduce_topics( doc_embeddings = np.array([topic_embeddings[topic_id] for topic_id in self.topics_]) corpus = Corpus( documents=docs, + media=images, + modality=Modality.IMAGE if images else Modality.TEXT, topics=np.array(self.topics_), - images=images, embeddings=doc_embeddings, ) @@ -2348,6 +2367,7 @@ def _extract_embeddings( documents = [documents] if images is not None and hasattr(self.embedding_model, "embed_images"): + documents = documents if any(documents) else [None] embeddings = self.embedding_model.embed(documents=documents, images=images, verbose=verbose) elif documents is not None: embeddings = self.embedding_model.embed_documents(documents, verbose=verbose) diff --git a/bertopic/_corpus.py b/bertopic/_corpus.py index bc7d596e..2f0e51ab 100644 --- a/bertopic/_corpus.py +++ b/bertopic/_corpus.py @@ -1,4 +1,5 @@ from dataclasses import dataclass, field +from enum import Enum import numpy as np from scipy.sparse import csr_matrix from collections import defaultdict @@ -6,16 +7,36 @@ from bertopic._topics import Topics +class Modality(str, Enum): + """What a row is, which determines whether its source lives in `documents` or `media`.""" + + TEXT = "text" + CODE = "code" + IMAGE = "image" + AUDIO = "audio" + VIDEO = "video" + + @dataclass class Corpus: - """Temporary container used to track the input and generated data during fitting.""" + """Temporary container used to track the input and generated data during fitting. + + A row's source lives in exactly one place, and `modality` says where: `TEXT` and + `CODE` rows are held in `documents` with no `media`, while `IMAGE`, `AUDIO` and + `VIDEO` rows are held in `media` with `documents` carrying their text surrogate. + A row may have both, which is how a captioned image is stored. + + Only `documents` feeds c-TF-IDF, so media rows without a surrogate yield blank + keywords rather than an error. + """ # Input data documents: list[str] | np.ndarray = field(default_factory=list) + media: list = field(default_factory=list) + modality: list[Modality] | Modality | None = None topics: np.ndarray | None = None probabilities: np.ndarray | None = None embeddings: np.ndarray | None = None - images: list[str] | None = None timestamps: list[str] | list[int] | np.ndarray | None = None classes: list[str] | list[int] | np.ndarray | None = None @@ -30,9 +51,6 @@ class Corpus: _zeroshot_labels: list[str] = field(default_factory=list) def __post_init__(self): - if self.original_indices is None: - self.original_indices = np.arange(len(self.documents)) - # For inference where a single document is passed as a string if isinstance(self.documents, str): self.documents = [self.documents] @@ -40,6 +58,24 @@ def __post_init__(self): if isinstance(self.documents, np.ndarray): self.documents = self.documents.tolist() + if self.documents is None: + self.documents = [] + + if self.media is None: + self.media = [] + + # Every row appears in both channels, with an empty surrogate until one is made + nr_rows = max(len(self.documents), len(self.media)) + self.documents = self.documents or [""] * nr_rows + self.media = self.media or [None] * nr_rows + + # A single modality applies to every row + if not isinstance(self.modality, list): + self.modality = [self.modality or Modality.TEXT] * nr_rows + + if self.original_indices is None: + self.original_indices = np.arange(nr_rows) + if isinstance(self.classes, list): self.classes = np.array(self.classes) @@ -48,16 +84,21 @@ def __post_init__(self): check_documents_type(self.documents) check_embeddings_shape(self.embeddings, self.documents) + self._validate_length("media", self.media) + self._validate_length("modality", self.modality) + + # Later assignments are guarded too, now that every channel is normalised + self._initialized = True @property - def has_only_images(self) -> bool: - """Check whether only images are provided.""" - return self.images is not None and self.documents is None + def images(self) -> list: + """The image rows of `media`.""" + return [item for item, modality in zip(self.media, self.modality) if modality == Modality.IMAGE] @property - def has_documents(self) -> bool: - """Check whether documents are provided.""" - return self.documents is not None + def has_only_images(self) -> bool: + """Check whether there are image rows and no text to describe them.""" + return bool(self.images) and not any(self.documents) @property def has_outliers(self) -> bool: @@ -237,11 +278,12 @@ def sort_by_timestamps(self) -> "Corpus": sort_order = np.argsort(self.timestamps) - self.documents = [self.documents[i] for i in sort_order] + self.documents = [self.documents[index] for index in sort_order] self.topics = np.array(self.topics)[sort_order] if self.topics is not None else None self.probabilities = self.probabilities[sort_order] if self.probabilities is not None else None self.embeddings = self.embeddings[sort_order] if self.embeddings is not None else None - self.images = list(np.array(self.images)[sort_order]) if self.images is not None else None + self.media = [self.media[index] for index in sort_order] + self.modality = [self.modality[index] for index in sort_order] self.timestamps = self.timestamps[sort_order] self.classes = np.array(self.classes)[sort_order] if self.classes is not None else None self.umap_embeddings = ( @@ -266,15 +308,13 @@ def get_corpus_by_indices(self, indices: list[int]) -> "Corpus": [self.topics[index] for index in sorted_indices] if self.topics is not None else None ) selected_embeddings = self.embeddings[sorted_indices] if self.embeddings is not None else None - selected_images = ( - [self.images[index] for index in sorted_indices] if self.images is not None else None - ) selected_original_indices = [self.original_indices[index] for index in sorted_indices] return Corpus( documents=selected_documents, + media=[self.media[index] for index in sorted_indices], + modality=[self.modality[index] for index in sorted_indices], topics=selected_topics, embeddings=selected_embeddings, - images=selected_images, original_indices=selected_original_indices, _zeroshot_labels=self._zeroshot_labels, ) @@ -304,6 +344,8 @@ def get_topic(self, topic_id: int, nr_samples: int | None = None) -> "Corpus": return Corpus( documents=filtered_docs, + media=[self.media[index] for index in filtered_indices], + modality=[self.modality[index] for index in filtered_indices], topics=[topic_id] * len(filtered_docs), embeddings=filtered_embeddings, original_indices=filtered_indices, @@ -319,9 +361,11 @@ def __add__(self, other: "Corpus") -> "Corpus": # Get the sorting order to restore original order sort_order = np.argsort(combined_indices) - # Combine and reorder documents + # Combine and reorder both channels combined_documents = self.documents + other.documents - sorted_documents = [combined_documents[i] for i in sort_order] + sorted_documents = [combined_documents[index] for index in sort_order] + combined_media = self.media + other.media + combined_modality = self.modality + other.modality # Combine and reorder topics if other.has_zeroshot_labels: @@ -344,6 +388,8 @@ def __add__(self, other: "Corpus") -> "Corpus": return Corpus( documents=sorted_documents, + media=[combined_media[index] for index in sort_order], + modality=[combined_modality[index] for index in sort_order], topics=sorted_topics, embeddings=sorted_embeddings, original_indices=sorted_indices, @@ -361,8 +407,8 @@ def _validate_length(self, name: str, value) -> None: ) def __setattr__(self, name: str, value) -> None: - """Whenever we update embeddings, images, or topics, validate their length.""" - if name in ("embeddings", "images", "topics") and hasattr(self, "documents"): + """Whenever we update a per-row field after construction, validate its length.""" + if name in ("embeddings", "media", "modality", "topics") and getattr(self, "_initialized", False): self._validate_length(name, value) super().__setattr__(name, value) diff --git a/tests/test_corpus.py b/tests/test_corpus.py index 178986f2..fae715f4 100644 --- a/tests/test_corpus.py +++ b/tests/test_corpus.py @@ -1,8 +1,9 @@ """Contract tests for the Corpus container. `Corpus` is the value object that carries documents, embeddings, and assignments -through the pipeline. These tests pin down what it accepts, which matters because -its `__post_init__` validation is what currently rejects image-only input. +through the pipeline. Text and media are parallel channels: `documents` feeds +c-TF-IDF, `media` holds whatever a row is made of that is not text, and `modality` +says what `media` holds. These tests pin down that contract. """ import importlib.util @@ -10,7 +11,7 @@ import numpy as np import pytest -from bertopic._corpus import Corpus +from bertopic._corpus import Corpus, Modality def pillow_available(): @@ -57,19 +58,125 @@ def test_mismatched_topics_are_rejected(): @pytest.mark.skipif(not pillow_available(), reason="Pillow not available") -@pytest.mark.xfail( - strict=True, - reason="Bug 1: image-only input raises in __post_init__; fixed in unit 2", -) -def test_images_may_be_supplied_without_documents(image_paths): - """Multimodal input has no documents, which is the documented API for images. - - `docs/getting_started/multimodal/multimodal.md` calls - `fit_transform(documents=None, images=images)`, but `check_documents_type` rejects - `None` before any embedding happens, so `has_only_images` can never be True and - `_images_to_text` is unreachable. - """ - corpus = Corpus(documents=None, images=image_paths) +def test_media_may_be_supplied_without_documents(image_paths): + """Image-only input used to raise in `__post_init__` before any embedding ran.""" + corpus = Corpus(media=image_paths, modality=Modality.IMAGE) + assert len(corpus) == len(image_paths) + assert corpus.images == image_paths assert corpus.has_only_images - assert len(corpus.images) == len(image_paths) + + +def test_documents_default_to_the_text_modality(): + """Text rows are their own source, so nothing is stored twice.""" + corpus = Corpus(documents=["first", "second"]) + + assert corpus.modality == [Modality.TEXT, Modality.TEXT] + assert corpus.media == [None, None] + + +def test_a_single_modality_applies_to_every_row(): + """Most corpora are uniform, so a scalar modality is broadcast.""" + corpus = Corpus(media=["a.png", "b.png", "c.png"], modality=Modality.IMAGE) + + assert corpus.modality == [Modality.IMAGE] * 3 + + +def test_a_row_may_carry_both_text_and_media(): + """An image with a caption is one row with both channels populated.""" + corpus = Corpus(documents=["a cat", "a dog"], media=["cat.png", "dog.png"], modality=Modality.IMAGE) + + assert corpus.documents == ["a cat", "a dog"] + assert corpus.images == ["cat.png", "dog.png"] + assert not corpus.has_only_images + + +def test_rows_may_each_have_their_own_modality(): + """Unrelated sets of text and images can share a corpus.""" + corpus = Corpus( + documents=["some text", ""], + media=[None, "picture.png"], + modality=[Modality.TEXT, Modality.IMAGE], + ) + + assert corpus.images == ["picture.png"] + assert len(corpus) == 2 + + +def test_rows_without_text_keep_an_empty_document(): + """The text channel stays addressable until a representation model fills it in.""" + corpus = Corpus(media=["a.png", "b.png"], modality=Modality.IMAGE) + + assert corpus.documents == ["", ""] + + +def test_selecting_by_index_carries_both_channels(): + """Slicing keeps every row's media and modality alongside its text.""" + corpus = Corpus( + documents=["first", "second", "third"], + media=["a.png", "b.png", "c.png"], + modality=Modality.IMAGE, + ) + + selected = corpus.get_corpus_by_indices([0, 2]) + + assert selected.documents == ["first", "third"] + assert selected.media == ["a.png", "c.png"] + assert selected.modality == [Modality.IMAGE, Modality.IMAGE] + + +def test_selecting_by_topic_carries_both_channels(): + """`get_topic` used to drop image media entirely.""" + corpus = Corpus( + documents=["first", "second", "third"], + media=["a.png", "b.png", "c.png"], + modality=Modality.IMAGE, + topics=np.array([0, 1, 0]), + ) + + selected = corpus.get_topic(0) + + assert selected.media == ["a.png", "c.png"] + assert selected.images == ["a.png", "c.png"] + + +def test_combining_corpora_carries_both_channels(): + """Zero-shot recombines two corpora and must not lose media.""" + first = Corpus( + documents=["first"], + media=["a.png"], + modality=Modality.IMAGE, + topics=np.array([0]), + embeddings=np.zeros((1, 4)), + original_indices=[0], + ) + second = Corpus( + documents=["second"], + media=["b.png"], + modality=Modality.IMAGE, + topics=np.array([1]), + embeddings=np.ones((1, 4)), + original_indices=[1], + ) + + combined = first + second + + assert combined.media == ["a.png", "b.png"] + assert combined.modality == [Modality.IMAGE, Modality.IMAGE] + + +def test_mismatched_channels_are_rejected_at_construction(): + """A short channel would otherwise be silently truncated by every later zip.""" + with pytest.raises(ValueError): + Corpus(documents=["first", "second"], media=["only-one.png"], modality=Modality.IMAGE) + + with pytest.raises(ValueError): + Corpus(documents=["first", "second"], modality=[Modality.TEXT]) + + +def test_mismatched_media_is_rejected(): + """Content is a per-row field, so its length is guarded like the others.""" + corpus = Corpus(documents=["first", "second"]) + + with pytest.raises(ValueError): + corpus.media = ["only one"]