Skip to content
Merged
Show file tree
Hide file tree
Changes from 11 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
35 changes: 28 additions & 7 deletions mteb/abstasks/AbsTask.py
Original file line number Diff line number Diff line change
Expand Up @@ -408,25 +408,46 @@ def filter_languages(
def _add_main_score(self, scores: dict[HFSubset, ScoresDict]) -> None:
scores["main_score"] = scores[self.metadata.main_score]

def _upload_dataset_to_hub(self, repo_name: str, fields: list[str]) -> None:
def _upload_dataset_to_hub(
self, repo_name: str, fields: list[str] | dict[str, str]
) -> None:
if self.metadata.is_multilingual:
for config in self.metadata.eval_langs:
logger.info(f"Converting {config} of {self.metadata.name}")
sentences = {}
for split in self.dataset[config]:
sentences[split] = Dataset.from_dict(
{field: self.dataset[config][split][field] for field in fields}
)
if isinstance(fields, dict):
sentences[split] = Dataset.from_dict(
{
mapped_name: self.dataset[config][split][original_name]
for original_name, mapped_name in fields.items()
}
)
else:
sentences[split] = Dataset.from_dict(
{
field: self.dataset[config][split][field]
for field in fields
}
)
sentences = DatasetDict(sentences)
sentences.push_to_hub(
repo_name, config, commit_message=f"Add {config} dataset"
)
else:
sentences = {}
for split in self.dataset:
sentences[split] = Dataset.from_dict(
{field: self.dataset[split][field] for field in fields}
)
if isinstance(fields, dict):
sentences[split] = Dataset.from_dict(
{
mapped_name: self.dataset[split][original_name]
for original_name, mapped_name in fields.items()
}
)
else:
sentences[split] = Dataset.from_dict(
{field: self.dataset[split][field] for field in fields}
)
sentences = DatasetDict(sentences)
sentences.push_to_hub(repo_name, commit_message="Add dataset")

Expand Down
Comment thread
Samoed marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from ..evaluation.evaluators import (
logRegClassificationEvaluator,
)
from ..evaluation.evaluators.ClassificationEvaluator import AbsClassificationEvaluator
from ..load_results.task_results import HFSubset, ScoresDict
from .AbsTask import AbsTask

Expand All @@ -35,6 +36,16 @@ class ClassificationDescriptiveStatistics(DescriptiveStatistics):
min_labels_per_text: Minimum number of labels per text
average_label_per_text: Average number of labels per text
max_labels_per_text: Maximum number of labels per text

min_image_width: Minimum width of images
average_image_width: Average width of images
max_image_width: Maximum width of images

min_image_height: Minimum height of images
average_image_height: Average height of images
max_image_height: Maximum height of images


unique_labels: Number of unique labels
labels: dict of label frequencies
"""
Expand All @@ -51,11 +62,20 @@ class ClassificationDescriptiveStatistics(DescriptiveStatistics):
min_labels_per_text: int
average_label_per_text: float
max_labels_per_text: int

min_image_width: float | None
average_image_width: float | None
max_image_width: float | None

min_image_height: float | None
average_image_height: float | None
max_image_height: float | None

unique_labels: int
labels: dict[str, dict[str, int]]


class AbsTaskClassification(AbsTask):
class AbsTaskAnyClassification(AbsTask):
Comment thread
Samoed marked this conversation as resolved.
"""Abstract class for classification tasks
The similarity is computed between pairs and the results are ranked.

Expand All @@ -69,12 +89,14 @@ class AbsTaskClassification(AbsTask):

"""

evaluator = logRegClassificationEvaluator
abstask_prompt = "Classify user passages."
evaluator: type[AbsClassificationEvaluator] = logRegClassificationEvaluator
samples_per_label: int = 8
n_experiments: int = 10
k: int = 3
train_split = "train"
train_split: str = "train"
label_column_name: str = "label"
values_column_name: str = "text"
Comment thread
Samoed marked this conversation as resolved.
Outdated
abstask_prompt = "Classify user passages."
Comment thread
KennethEnevoldsen marked this conversation as resolved.

def evaluate(
self,
Expand Down Expand Up @@ -117,7 +139,7 @@ def evaluate(
def _evaluate_subset(
self,
model: Encoder,
dataset: DatasetDict | Dataset,
dataset: DatasetDict,
hf_split: str,
hf_subset: str,
encode_kwargs: dict[str, Any],
Expand All @@ -128,6 +150,10 @@ def _evaluate_subset(
params = {"k": self.k}
params.update(kwargs)

is_image = False
if "image" in self.metadata.modalities:
is_image = True

scores = []
test_cache, idxs = (
None,
Expand All @@ -140,13 +166,15 @@ def _evaluate_subset(
# Bootstrap `self.samples_per_label` samples per label for each split
train_dataset, idxs = self._undersample_data(
train_split,
self.samples_per_label,
idxs,
)

evaluator = self.evaluator(
train_dataset,
eval_split,
self.values_column_name,
self.label_column_name,
is_image,
task_metadata=self.metadata,
hf_split=hf_split,
hf_subset=hf_subset,
Expand All @@ -164,13 +192,12 @@ def _evaluate_subset(
return avg_scores

def _undersample_data(
self, dataset: Dataset, samples_per_label: int, idxs=None
self, dataset: Dataset, idxs: list[int] | None = None
) -> tuple[Dataset, list[int]]:
"""Undersample data to have `samples_per_label` samples of each label.

Args:
dataset: Hugging Face `datasets.Dataset` containing "text" and "label".
samples_per_label: Number of samples per label to retain.
idxs: Optional indices to shuffle and sample from.

Returns:
Expand All @@ -187,8 +214,8 @@ def _undersample_data(
sampled_idxs = []

for i in idxs:
label = dataset[i]["label"]
if label_counter[label] < samples_per_label:
label = dataset[i][self.label_column_name]
if label_counter[label] < self.samples_per_label:
sampled_idxs.append(i)
label_counter[label] += 1

Expand All @@ -199,26 +226,50 @@ def _calculate_metrics_from_split(
) -> ClassificationDescriptiveStatistics:
train_text = []
if hf_subset:
text = self.dataset[hf_subset][split]["text"]
label = self.dataset[hf_subset][split]["label"]
if split != "train":
train_text = self.dataset[hf_subset]["train"]["text"]
values = self.dataset[hf_subset][split][self.values_column_name]
label = self.dataset[hf_subset][split][self.label_column_name]
if split != self.train_split:
train_text = self.dataset[hf_subset][self.train_split][
self.values_column_name
]
elif compute_overall:
text = []
values = []
Comment thread
Samoed marked this conversation as resolved.
Outdated
label = []
for hf_subset in self.metadata.eval_langs:
text.extend(self.dataset[hf_subset][split]["text"])
label.extend(self.dataset[hf_subset][split]["label"])
if split != "train":
train_text.extend(self.dataset[hf_subset]["train"]["text"])
values.extend(self.dataset[hf_subset][split][self.values_column_name])
label.extend(self.dataset[hf_subset][split][self.label_column_name])
if split != self.train_split:
train_text.extend(
self.dataset[hf_subset][self.train_split][
self.values_column_name
]
)
else:
text = self.dataset[split]["text"]
label = self.dataset[split]["label"]
if split != "train":
train_text = self.dataset["train"]["text"]
values = self.dataset[split][self.values_column_name]
label = self.dataset[split][self.label_column_name]
if split != self.train_split:
train_text = self.dataset[self.train_split][self.values_column_name]

total_text_len = 0
text_len = None
img_widths, img_heights = None, None
num_texts_in_train = None

if "image" in self.metadata.modalities:
img_widths, img_heights = [], []
for img in values:
width, height = img.size # type: ignore
img_heights.append(height)
img_widths.append(width)
else:
Comment thread
Samoed marked this conversation as resolved.
Outdated
text_len = [len(t) for t in values]
total_text_len = sum(text_len)
num_texts_in_train = (
len(set(values) & set(train_text))
if split != self.train_split
else None
)

text_len = [len(t) for t in text]
total_text_len = sum(text_len)
if isinstance(label[0], int):
label_len = [1] * len(label)
total_label_len = len(label)
Expand All @@ -230,18 +281,32 @@ def _calculate_metrics_from_split(
total_labels = []
for l in label:
total_labels.extend(l if len(l) > 0 else [None])

label_count = Counter(total_labels)
num_texts_in_train = (
len(set(text) & set(train_text)) if split != "train" else None
)

return ClassificationDescriptiveStatistics(
num_samples=len(text),
num_samples=len(values),
# text
number_of_characters=total_text_len,
number_texts_intersect_with_train=num_texts_in_train,
min_text_length=min(text_len),
average_text_length=total_text_len / len(text),
max_text_length=max(text_len),
unique_texts=len(set(text)),
number_texts_intersect_with_train=num_texts_in_train
if num_texts_in_train
else None,
min_text_length=min(text_len) if text_len else None,
average_text_length=total_text_len / len(values) if text_len else None,
max_text_length=max(text_len) if text_len else None,
unique_texts=len(set(values)) if text_len else None,
# image
min_image_width=min(img_widths) if img_widths else None,
average_image_width=sum(img_widths) / len(img_widths)
if img_widths
else None,
max_image_width=max(img_widths) if img_widths else None,
min_image_height=min(img_heights) if img_heights else None,
average_image_height=sum(img_heights) / len(img_heights)
if img_heights
else None,
max_image_height=max(img_heights) if img_heights else None,
# labels
min_labels_per_text=min(label_len),
average_label_per_text=total_label_len / len(label),
max_labels_per_text=max(label_len),
Expand All @@ -255,4 +320,10 @@ def _calculate_metrics_from_split(
)

def _push_dataset_to_hub(self, repo_name: str) -> None:
self._upload_dataset_to_hub(repo_name, ["text", "label"])
self._upload_dataset_to_hub(
repo_name,
[
self.values_column_name,
self.label_column_name,
],
)
4 changes: 2 additions & 2 deletions mteb/abstasks/AbsTaskMultilabelClassification.py
Comment thread
Samoed marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from mteb.encoder_interface import Encoder

from ..load_results.task_results import ScoresDict
from .AbsTaskClassification import AbsTaskClassification
from .AbsTaskAnyClassification import AbsTaskAnyClassification

logger = logging.getLogger(__name__)

Expand All @@ -41,7 +41,7 @@ def evaluate_classifier(
}


class AbsTaskMultilabelClassification(AbsTaskClassification):
class AbsTaskMultilabelClassification(AbsTaskAnyClassification):
"""Abstract class for multioutput classification tasks
The similarity is computed between pairs and the results are ranked.

Expand Down
Loading