From b822fb0ffa71c1c039fc6b45cd68725c280f58db Mon Sep 17 00:00:00 2001 From: dajiaohuang Date: Thu, 24 Sep 2026 03:55:48 +0800 Subject: [PATCH] Default unset focal loss reduction --- gliclass/model.py | 3 ++- tests/test_loss_functions.py | 25 +++++++++++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/gliclass/model.py b/gliclass/model.py index 8d8702b..0757381 100644 --- a/gliclass/model.py +++ b/gliclass/model.py @@ -329,12 +329,13 @@ def get_loss(self, logits, labels, classes_embedding=None, classes_embedding_mas loss_fct = nn.CrossEntropyLoss() loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1)) elif self.config.problem_type == "multi_label_classification": + reduction = self.config.focal_loss_reduction or "none" all_losses = focal_loss_with_logits( logits, labels, self.config.focal_loss_alpha, self.config.focal_loss_gamma, - self.config.focal_loss_reduction, + reduction, ) if classes_embedding_mask is not None: all_losses = all_losses * classes_embedding_mask.float() diff --git a/tests/test_loss_functions.py b/tests/test_loss_functions.py index 1bc3437..e0d5301 100644 --- a/tests/test_loss_functions.py +++ b/tests/test_loss_functions.py @@ -1,8 +1,11 @@ """Tests for gliclass.loss_functions module.""" +from types import SimpleNamespace + import pytest import torch +from gliclass.model import GLiClassBaseModel from gliclass.loss_functions import sequence_contrastive_loss, focal_loss_with_logits @@ -184,3 +187,25 @@ def test_all_ones_targets(self): assert not torch.isnan(loss) assert loss >= 0 + + +def test_base_model_defaults_missing_focal_reduction_to_none_mode(): + model = SimpleNamespace( + config=SimpleNamespace( + problem_type="multi_label_classification", + focal_loss_alpha=0.5, + focal_loss_gamma=2, + focal_loss_reduction=None, + contrastive_loss_coef=0, + ), + num_labels=2, + ) + + loss = GLiClassBaseModel.get_loss( + model, + logits=torch.tensor([[0.0, 1.0]]), + labels=torch.tensor([[0.0, 1.0]]), + ) + + assert loss.dim() == 0 + assert torch.isfinite(loss)