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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 26 additions & 6 deletions bertopic/_bertopic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
)

Expand Down Expand Up @@ -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,
)

Expand Down Expand Up @@ -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)
Expand Down
88 changes: 67 additions & 21 deletions bertopic/_corpus.py
Original file line number Diff line number Diff line change
@@ -1,21 +1,42 @@
from dataclasses import dataclass, field
from enum import Enum
import numpy as np
from scipy.sparse import csr_matrix
from collections import defaultdict

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

Expand All @@ -30,16 +51,31 @@ 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]

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)

Expand All @@ -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:
Expand Down Expand Up @@ -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 = (
Expand All @@ -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,
)
Expand Down Expand Up @@ -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,
Expand All @@ -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:
Expand All @@ -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,
Expand All @@ -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)

Expand Down
Loading