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
75 changes: 73 additions & 2 deletions model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,71 @@ 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.

: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 = _focal_modulation(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 (_focal_modulation(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.

Expand Down Expand Up @@ -271,6 +336,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.

Expand Down Expand Up @@ -309,8 +375,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)
Expand All @@ -325,10 +396,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(
Expand Down
70 changes: 69 additions & 1 deletion tests/test_trainable.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -894,6 +894,74 @@ 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)
Comment thread
stephantul marked this conversation as resolved.


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


@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
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
Expand Down
Loading