Commit ·
a770449
1
Parent(s): 251e9cd
Update LSTM gradient clipping
Browse files
knowledge_distillation.py
CHANGED
|
@@ -23,6 +23,7 @@ class DistillationTrainer:
|
|
| 23 |
alpha=0.5, # Weight for distillation loss vs. regular loss
|
| 24 |
lr=0.001,
|
| 25 |
weight_decay=1e-5,
|
|
|
|
| 26 |
device=None
|
| 27 |
):
|
| 28 |
self.teacher_model = teacher_model
|
|
@@ -32,6 +33,7 @@ class DistillationTrainer:
|
|
| 32 |
self.test_loader = test_loader
|
| 33 |
self.temperature = temperature
|
| 34 |
self.alpha = alpha
|
|
|
|
| 35 |
|
| 36 |
self.device = device if device else torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 37 |
logger.info(f"Using device: {self.device}")
|
|
@@ -137,7 +139,7 @@ class DistillationTrainer:
|
|
| 137 |
# Backward and optimize
|
| 138 |
self.optimizer.zero_grad()
|
| 139 |
loss.backward()
|
| 140 |
-
torch.nn.utils.clip_grad_norm_(self.student_model.parameters(),
|
| 141 |
self.optimizer.step()
|
| 142 |
|
| 143 |
train_loss += loss.item()
|
|
|
|
| 23 |
alpha=0.5, # Weight for distillation loss vs. regular loss
|
| 24 |
lr=0.001,
|
| 25 |
weight_decay=1e-5,
|
| 26 |
+
max_grad_norm=1.0,
|
| 27 |
device=None
|
| 28 |
):
|
| 29 |
self.teacher_model = teacher_model
|
|
|
|
| 33 |
self.test_loader = test_loader
|
| 34 |
self.temperature = temperature
|
| 35 |
self.alpha = alpha
|
| 36 |
+
self.max_grad_norm = max_grad_norm
|
| 37 |
|
| 38 |
self.device = device if device else torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 39 |
logger.info(f"Using device: {self.device}")
|
|
|
|
| 139 |
# Backward and optimize
|
| 140 |
self.optimizer.zero_grad()
|
| 141 |
loss.backward()
|
| 142 |
+
torch.nn.utils.clip_grad_norm_(self.student_model.parameters(), self.max_grad_norm)
|
| 143 |
self.optimizer.step()
|
| 144 |
|
| 145 |
train_loss += loss.item()
|