From 24b907953137fa733db42a10c1d0022d0c763c1c Mon Sep 17 00:00:00 2001 From: stephantul Date: Thu, 8 Oct 2026 11:47:06 +0200 Subject: [PATCH 1/2] feat: add focal loss --- model2vec/train/classifier.py | 63 +++++++++++++++++++++++++++++++++-- tests/test_trainable.py | 46 ++++++++++++++++++++++++- 2 files changed, 106 insertions(+), 3 deletions(-) diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index fb0d3f6..5e579dd 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -46,6 +46,59 @@ def _multilabel_classifier_metrics(head_out: torch.Tensor, y: torch.Tensor, loss return {"loss": loss.item(), "accuracy": accuracy} +class FocalLoss(nn.Module): + def __init__(self, gamma: float = 0.0, weight: torch.Tensor | None = None) -> None: + """Initialize the focal loss for single-label classification. + + :param gamma: The focusing parameter. If this is 0.0, the loss is equal to cross-entropy. + :param weight: The weight of each class, or None to weight all classes equally. + """ + super().__init__() + self.gamma = gamma + self.weight: torch.Tensor | None + self.register_buffer("weight", weight) + + def forward(self, head_out: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + """Returns the focal loss, averaged like weighted cross-entropy. + + :param head_out: The logits. + :param y: The class indices. + :return: The mean loss. + """ + nll = -torch.log_softmax(head_out, dim=1).gather(1, y[:, None]).squeeze(1) + loss = (1 - torch.exp(-nll)) ** self.gamma * nll + if self.weight is None: + return loss.mean() + sample_weight = self.weight[y] + return (loss * sample_weight).sum() / sample_weight.sum() + + +class BinaryFocalLoss(nn.Module): + def __init__(self, gamma: float = 0.0, pos_weight: torch.Tensor | None = None) -> None: + """Initialize the focal loss for multi-label classification. + + :param gamma: The focusing parameter. If this is 0.0, the loss is equal to binary cross-entropy. + :param pos_weight: The weight of the positive examples of each class, or None to weight them equally. + """ + super().__init__() + self.gamma = gamma + self.pos_weight: torch.Tensor | None + self.register_buffer("pos_weight", pos_weight) + + def forward(self, head_out: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + """Returns the focal loss, averaged over all samples and classes. + + :param head_out: The logits. + :param y: The multi-hot targets. + :return: The mean loss. + """ + loss = nn.functional.binary_cross_entropy_with_logits(head_out, y, pos_weight=self.pos_weight, reduction="none") + if self.gamma == 0: + return loss.mean() + nll = nn.functional.binary_cross_entropy_with_logits(head_out, y, reduction="none") + return ((1 - torch.exp(-nll)) ** self.gamma * loss).mean() + + def _read_labels(y: LabelType, name: str) -> tuple[bool, Counter]: """Determine whether labels are multi-label, and count the number of times each class occurs. @@ -271,6 +324,7 @@ def fit( validation_steps: int | None = None, random_seed: int = DEFAULT_RANDOM_SEED, token_dropout: float = 0.0, + focal_gamma: float = 0.0, ) -> StaticModelForClassification: """Fit a model. @@ -309,8 +363,13 @@ def fit( :param random_seed: The random seed to use. Defaults to 42. :param token_dropout: The fraction of tokens to randomly drop from each training sample. Has no effect during validation. Must be in the range [0, 1). + :param focal_gamma: The gamma of the focal loss. If this is 0.0, (binary) cross-entropy is used. + Must be non-negative. :return: The fitted model. + :raises ValueError: If `focal_gamma` is negative. """ + if focal_gamma < 0: + raise ValueError(f"focal_gamma must be non-negative, got {focal_gamma}.") seed_everything(random_seed) logger.info("Re-initializing model.") self._check_inputs(X=X, y=y, X_val=X_val, y_val=y_val) @@ -325,10 +384,10 @@ def fit( ) if self.multilabel: - loss_function: nn.Module = nn.BCEWithLogitsLoss(pos_weight=resolved_class_weight) + loss_function: nn.Module = BinaryFocalLoss(gamma=focal_gamma, pos_weight=resolved_class_weight) compute_metrics = _multilabel_classifier_metrics else: - loss_function = nn.CrossEntropyLoss(weight=resolved_class_weight) + loss_function = FocalLoss(gamma=focal_gamma, weight=resolved_class_weight) compute_metrics = _classifier_metrics self._train( diff --git a/tests/test_trainable.py b/tests/test_trainable.py index ca1cc35..e08f0f3 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -19,7 +19,7 @@ from model2vec.model import StaticModel from model2vec.train import StaticModelForClassification from model2vec.train.base import BaseFinetuneable -from model2vec.train.classifier import _read_labels +from model2vec.train.classifier import BinaryFocalLoss, FocalLoss, _read_labels from model2vec.train.dataset import ( ColumnRows, PairDataset, @@ -894,6 +894,50 @@ def test_y_val_none() -> None: model.fit(X, y, X_val=None, y_val=None) +@pytest.mark.parametrize("weight", [None, torch.tensor([1.0, 2.0, 0.5])]) +def test_focal_loss_with_zero_gamma_is_cross_entropy(weight: torch.Tensor | None) -> None: + """With gamma 0, the focal loss equals (weighted) cross-entropy.""" + torch.random.manual_seed(42) + logits = torch.randn(8, 3) + y = torch.randint(0, 3, (8,)) + expected = nn.CrossEntropyLoss(weight=weight)(logits, y) + assert torch.allclose(FocalLoss(gamma=0.0, weight=weight)(logits, y), expected) + + +@pytest.mark.parametrize("pos_weight", [None, torch.tensor([1.0, 2.0, 0.5])]) +def test_binary_focal_loss_with_zero_gamma_is_binary_cross_entropy(pos_weight: torch.Tensor | None) -> None: + """With gamma 0, the binary focal loss equals (weighted) binary cross-entropy.""" + torch.random.manual_seed(42) + logits = torch.randn(8, 3) + y = torch.randint(0, 2, (8, 3)).float() + expected = nn.BCEWithLogitsLoss(pos_weight=pos_weight)(logits, y) + assert torch.allclose(BinaryFocalLoss(gamma=0.0, pos_weight=pos_weight)(logits, y), expected) + + +def test_focal_loss_downweights_easy_examples() -> None: + """A positive gamma shrinks the loss of a confident correct prediction more than that of a wrong one.""" + logits = torch.tensor([[4.0, 0.0], [0.0, 4.0]]) + y = torch.tensor([0, 0]) + ce = nn.CrossEntropyLoss(reduction="none")(logits, y) + easy = FocalLoss(gamma=2.0)(logits[:1], y[:1]) + hard = FocalLoss(gamma=2.0)(logits[1:], y[1:]) + assert easy / ce[0] < hard / ce[1] < 1 + + +def test_fit_with_focal_gamma() -> None: + """fit() trains with a focal loss, and rejects a negative gamma.""" + tokenizer = AutoTokenizer.from_pretrained("tests/data/test_tokenizer").backend_tokenizer + torch.random.manual_seed(42) + vectors_torched = torch.randn(len(tokenizer.get_vocab()), 12) + model = StaticModelForClassification(vectors=vectors_torched, tokenizer=tokenizer, hidden_dim=12).to("cpu") + + X = ["dog", "cat"] + with pytest.raises(ValueError): + model.fit(X, ["0", "1"], focal_gamma=-1.0, max_epochs=1) + model.fit(X, ["0", "1"], focal_gamma=2.0, class_weight="balanced", max_epochs=1) + model.fit(X, [["0"], ["0", "1"]], focal_gamma=2.0, class_weight="balanced", max_epochs=1) + + def test_class_weight() -> None: """Test the class weight function.""" tokenizer = AutoTokenizer.from_pretrained("tests/data/test_tokenizer").backend_tokenizer From 5a7052034512abaa182b8b839b4313ef49033703 Mon Sep 17 00:00:00 2001 From: stephantul Date: Thu, 8 Oct 2026 12:06:14 +0200 Subject: [PATCH 2/2] fix for 0 --- model2vec/train/classifier.py | 16 ++++++++++++++-- tests/test_trainable.py | 24 ++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 2 deletions(-) diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index 5e579dd..322fa09 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -46,6 +46,18 @@ def _multilabel_classifier_metrics(head_out: torch.Tensor, y: torch.Tensor, loss return {"loss": loss.item(), "accuracy": accuracy} +def _focal_modulation(nll: torch.Tensor, gamma: float) -> torch.Tensor: + """Compute the focal modulating factor (1 - p) ** gamma from the negative log-likelihood. + + :param nll: The negative log-likelihood of the target. + :param gamma: The focusing parameter. + :return: The modulating factor. + """ + min_value = torch.finfo(nll.dtype).tiny + one_minus_p = (-torch.expm1(-nll)).clamp(min=min_value) + return one_minus_p**gamma + + class FocalLoss(nn.Module): def __init__(self, gamma: float = 0.0, weight: torch.Tensor | None = None) -> None: """Initialize the focal loss for single-label classification. @@ -66,7 +78,7 @@ def forward(self, head_out: torch.Tensor, y: torch.Tensor) -> torch.Tensor: :return: The mean loss. """ nll = -torch.log_softmax(head_out, dim=1).gather(1, y[:, None]).squeeze(1) - loss = (1 - torch.exp(-nll)) ** self.gamma * nll + loss = _focal_modulation(nll, self.gamma) * nll if self.weight is None: return loss.mean() sample_weight = self.weight[y] @@ -96,7 +108,7 @@ def forward(self, head_out: torch.Tensor, y: torch.Tensor) -> torch.Tensor: if self.gamma == 0: return loss.mean() nll = nn.functional.binary_cross_entropy_with_logits(head_out, y, reduction="none") - return ((1 - torch.exp(-nll)) ** self.gamma * loss).mean() + return (_focal_modulation(nll, self.gamma) * loss).mean() def _read_labels(y: LabelType, name: str) -> tuple[bool, Counter]: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index e08f0f3..00eb7a4 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -924,6 +924,30 @@ def test_focal_loss_downweights_easy_examples() -> None: assert easy / ce[0] < hard / ce[1] < 1 +@pytest.mark.parametrize("gamma", [0.0, 0.5, 2.0]) +def test_focal_loss_backward_with_confident_predictions(gamma: float) -> None: + """The focal loss has finite gradients when a prediction is confidently correct.""" + logits = torch.tensor([[100.0, 0.0], [0.0, 1.0]], requires_grad=True) + y = torch.tensor([0, 1]) + loss = FocalLoss(gamma=gamma, weight=torch.tensor([1.0, 2.0]))(logits, y) + loss.backward() + assert torch.isfinite(loss) + assert logits.grad is not None + assert torch.isfinite(logits.grad).all() + + +@pytest.mark.parametrize("gamma", [0.0, 0.5, 2.0]) +def test_binary_focal_loss_backward_with_confident_predictions(gamma: float) -> None: + """The binary focal loss has finite gradients when a prediction is confidently correct.""" + logits = torch.tensor([[100.0, -100.0], [0.5, 0.0]], requires_grad=True) + y = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) + loss = BinaryFocalLoss(gamma=gamma, pos_weight=torch.tensor([1.0, 2.0]))(logits, y) + loss.backward() + assert torch.isfinite(loss) + assert logits.grad is not None + assert torch.isfinite(logits.grad).all() + + def test_fit_with_focal_gamma() -> None: """fit() trains with a focal loss, and rejects a negative gamma.""" tokenizer = AutoTokenizer.from_pretrained("tests/data/test_tokenizer").backend_tokenizer