PyTorch中回调函数触发检测方法及BERT微调时触发ReduceLROnPlateau后解冻层的实现方案
Alright, let's break down your two questions one by one—super common scenarios when working with PyTorch and transformers like BERT!
PyTorch doesn’t have a built-in callback tracking system like Keras, but there are simple, effective ways to monitor when a callback fires:
Custom Callback with a Trigger Flag: The easiest approach is to extend or wrap existing scheduler/callback classes and add a boolean flag that gets flipped when the trigger condition is met. For example, for
ReduceLROnPlateau:import torch.optim.lr_scheduler as lr_scheduler class TrackableReduceLROnPlateau(lr_scheduler.ReduceLROnPlateau): def __init__(self, optimizer, *args, **kwargs): super().__init__(optimizer, *args, **kwargs) self.triggered = False # Flag to track activation def step(self, metrics, epoch=None): prev_lr = self.optimizer.param_groups[0]['lr'] super().step(metrics, epoch) current_lr = self.optimizer.param_groups[0]['lr'] # If learning rate changed, the callback was triggered self.triggered = current_lr != prev_lrAfter running
scheduler.step(val_loss)each epoch, you can checkscheduler.triggeredto confirm if it activated.Explicit Tracking in Custom Callbacks: For callbacks you write from scratch, just add a flag directly in the logic that runs on trigger. For example:
class CustomEarlyStopping: def __init__(self, patience=3): self.patience = patience self.best_loss = float('inf') self.counter = 0 self.triggered = False def step(self, val_loss): if val_loss < self.best_loss: self.best_loss = val_loss self.counter = 0 self.triggered = False else: self.counter += 1 if self.counter >= self.patience: self.triggered = TruePyTorch Lightning Shortcut (if using it): If you’re using Lightning, you can override
on_scheduler_stepin a custom callback to directly check if the scheduler triggered—no need for manual flagging.
This is a smart transfer learning strategy to preserve pre-trained weights while warming up the classifier. Here’s a step-by-step implementation:
Step 1: Freeze initial BERT layers
Start by freezing the base BERT model and only train the classifier head:
from transformers import BertForSequenceClassification model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2) # Freeze all base BERT layers (adjust to freeze only specific layers if needed) for param in model.bert.parameters(): param.requires_grad = False # Initialize optimizer to only update the classifier head optimizer = torch.optim.Adam(model.classifier.parameters(), lr=1e-4)
Step 2: Use the trackable scheduler from Question 1
Use our custom TrackableReduceLROnPlateau to detect when the learning rate is reduced:
scheduler = TrackableReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=3)
Step 3: Update training loop to unfreeze on trigger
In your training loop, check after each scheduler step if it triggered. If so, unfreeze the BERT layers and update the optimizer to include all parameters:
unfrozen = False # Prevent repeated unfreezing for epoch in range(num_epochs): # Training phase model.train() # ... your training steps (forward pass, loss calculation, optimizer step) ... # Validation phase model.eval() val_loss = ... # Compute your validation loss # Run scheduler step scheduler.step(val_loss) # Check if scheduler triggered and layers are still frozen if scheduler.triggered and not unfrozen: print("Scheduler activated—unfreezing BERT layers!") # Unfreeze all BERT layers for param in model.bert.parameters(): param.requires_grad = True # Update optimizer to train all parameters (use a smaller LR for fine-tuning) optimizer = torch.optim.Adam(model.parameters(), lr=1e-5) # Reset scheduler with the new optimizer (optional but recommended) scheduler = TrackableReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=3) unfrozen = True
Key Notes:
- Lower LR for Fine-Tuning: When unfreezing, use a smaller learning rate (like 1e-5) to avoid overwriting the pre-trained BERT weights.
- Partial Unfreezing: Instead of unfreezing all layers, you can unfreeze only the last 2-3 encoder layers for more controlled fine-tuning.
- Avoid Reinitializing Optimizer: If you don’t want to reinitialize the optimizer, you can modify its
param_groupsto add the now-trainable parameters, but reinitializing is simpler for most use cases.
内容的提问来源于stack exchange,提问作者Winvoker

