Keras中1700+类别多标签分类的优化器、损失函数及指标咨询
Hey there! Let's tackle your multi-label protein classification problem—9800 features, 1700 classes, 4-5 labels per sample—with practical recommendations for optimizers, loss functions, and a custom tensor-based accuracy metric.
Given your high-dimensional feature space and large number of classes, here are my top picks:
- AdamW: This is my go-to for most large-scale multi-label tasks. Unlike vanilla Adam, it decouples weight decay from gradient updates, which is critical for preventing overfitting when you have thousands of features and classes. Start with a learning rate of
1e-4to1e-3, and pair it with a weight decay of1e-5to1e-4—adjust based on your validation performance. - Cosine Annealing Learning Rate Scheduler: Even if your initial accuracy is high, scheduling the learning rate to gradually decrease helps the model converge to a more robust minimum. Combine this with AdamW using
torch.optim.lr_scheduler.CosineAnnealingLRfor best results. - SGD with Momentum (Backup Option): If you prefer a more stable optimizer with better generalization over long training runs, use SGD with momentum (0.9) and a learning rate of
1e-2, paired with weight decay. But note that AdamW will likely converge faster for your high-dimensional data.
For multi-label classification where each sample can belong to multiple classes, these losses work best:
- BCEWithLogitsLoss: The standard choice here. It combines sigmoid activation and binary cross-entropy in one step (numerically more stable than separate sigmoid + BCE). Since each class is treated as an independent binary classification problem, this aligns perfectly with your 4-5 labels per sample setup.
- Focal Loss: If you notice class imbalance (common in protein datasets where some functional classes are rare), Focal Loss reduces the weight of easy-to-classify samples, forcing the model to focus on harder, underrepresented classes. You can implement it by modifying BCEWithLogitsLoss with a gamma parameter (typically 2.0).
- Label Smoothing: Even with high accuracy, adding mild label smoothing (ε=0.1) can improve generalization by preventing the model from becoming overly confident in its predictions. You can integrate this into your BCE loss by adjusting the target labels slightly (e.g., 1 → 1-ε, 0 → ε).
Since you're working with tensors, here's an efficient PyTorch implementation of multi-label accuracy metrics. We'll cover two common variants: subset accuracy (strict match between predicted and true labels) and sample-level average accuracy (fraction of true labels correctly predicted per sample, averaged across all samples).
import torch def multi_label_accuracy(preds, targets, threshold=0.5): """ Computes multi-label accuracy metrics using vectorized tensor operations. Args: preds: Model output logits (shape: [batch_size, num_classes]) targets: Ground truth labels (shape: [batch_size, num_classes], 0/1 format) threshold: Threshold for converting logits to binary predictions Returns: subset_acc: Fraction of samples where predicted labels exactly match true labels sample_avg_acc: Average fraction of true labels correctly predicted per sample """ # Convert logits to binary predictions pred_labels = torch.sigmoid(preds) > threshold # Calculate subset accuracy (strict full match) full_match = torch.all(pred_labels == targets, dim=1) subset_acc = torch.mean(full_match.float()) # Calculate sample-level average accuracy (correct true labels / total true labels per sample) true_positives = torch.sum(pred_labels & targets, dim=1) total_true_labels = torch.sum(targets, dim=1) # Avoid division by zero for samples with no true labels (if any) sample_acc = true_positives / torch.clamp(total_true_labels, min=1) sample_avg_acc = torch.mean(sample_acc.float()) return subset_acc.item(), sample_avg_acc.item() # Example usage with dummy data batch_preds = torch.randn(32, 1700) # Batch of 32 samples, 1700 classes batch_targets = torch.randint(0, 2, (32, 1700)) # Random binary target labels subset_acc, sample_avg_acc = multi_label_accuracy(batch_preds, batch_targets) print(f"Subset Accuracy: {subset_acc:.4f}, Sample Average Accuracy: {sample_avg_acc:.4f}")
This implementation uses vectorized operations (no loops) for efficiency, which is crucial for handling your large dataset. Feel free to adjust the threshold based on your validation data—you might want to tune it to optimize for your specific task (e.g., higher threshold if you need fewer false positives).
内容的提问来源于stack exchange,提问作者Sarthak Mishra

