多分类不平衡NLP任务:PyTorch中Focal Loss结合类别权重的疑问
Hey there! Let's break down your questions one by one—dealing with imbalanced 4-class NLP tasks using BERT is a tricky but super common challenge, so I get where you're coming from.
1. Should we combine class weights with Focal Loss?
Absolutely. Here's why the two work great together:
- Focal Loss targets hard-to-classify samples by downweighting easy, well-predicted examples (via the
gammaparameter). This stops the model from just leaning into majority-class samples it already gets right. - Class weights (from
compute_class_weight) directly fix class imbalance by assigning higher weights to underrepresented classes, making the model care more about their samples.
These mechanisms complement each other perfectly—using both usually yields better results than either alone when your data is severely imbalanced.
2. Can we use the weight parameter in nn.CrossEntropyLoss() for this?
First a critical note: your model outputs nn.LogSoftmax(dim=1), but nn.CrossEntropyLoss() already includes a built-in LogSoftmax layer. Using it here would apply LogSoftmax twice, which is incorrect. You should use nn.NLLLoss() instead (which is what you were using originally with class weights).
That said: yes, you could pass class weights to nn.CrossEntropyLoss() if your model output raw logits (no LogSoftmax), but in your case, since you're using LogSoftmax in the model, you need to pair it with NLLLoss(weight=...) for the base loss calculation in your Focal Loss.
3. Correct implementation of Focal Loss with class weights
Your current Focal Loss code has two key issues:
- It uses
CrossEntropyLoss, which conflicts with your model's LogSoftmax output. - The
alphaparameter is a single global value, not class-specific (which is what you need for proper class weighting).
Here's the fixed implementation that integrates class weights properly:
import torch import torch.nn as nn from sklearn.utils.class_weight import compute_class_weight import numpy as np class FocalLossWithClassWeights(nn.Module): def __init__(self, class_weights, gamma=2.0, reduce=True): super().__init__() # Convert scikit-learn's class weights to a tensor self.class_weights = torch.tensor(class_weights, dtype=torch.float) self.gamma = gamma self.reduce = reduce def forward(self, inputs, targets): # inputs = model's LogSoftmax outputs (shape: [batch_size, num_classes]) # targets = ground-truth class indices (shape: [batch_size]) # Get log probability for the correct class log_pt = inputs[range(len(targets)), targets] # Convert log probability to actual probability pt = torch.exp(log_pt) # Fetch the weight for each sample's true class sample_weights = self.class_weights[targets].to(inputs.device) # Calculate focal loss per sample focal_loss = -sample_weights * (1 - pt)**self.gamma * log_pt if self.reduce: return torch.mean(focal_loss) else: return focal_loss # How to initialize it in your code: class_weights = compute_class_weight('balanced', np.unique(train_labels), train_labels) focal_loss = FocalLossWithClassWeights(class_weights=class_weights, gamma=2.0)
Key improvements:
- Matches your model's LogSoftmax output with an
NLLLoss-style calculation (no redundant LogSoftmax application). - Incorporates class-specific weights directly into the focal loss formula, so each sample's loss is scaled by its class's weight.
- Calculates
ptcorrectly as the model's predicted probability for the true class (no need to compute BCE loss first).
Quick tip:
The code automatically moves class weights to the same device (CPU/GPU) as your model inputs, so you don't have to handle that manually during training.
内容的提问来源于stack exchange,提问作者Deshwal

