Commit ·
ca56e1c
1
Parent(s): 01b0ca9
Save vocab size with model_state_dict
Browse files- inference_lstm.py +5 -0
- knowledge_distillation.py +4 -1
inference_lstm.py
CHANGED
|
@@ -51,6 +51,11 @@ if __name__ == "__main__":
|
|
| 51 |
output_dim=args.num_classes)
|
| 52 |
|
| 53 |
model_state = torch.load(args.model_path)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
if 'model_state_dict' in model_state:
|
| 55 |
model.load_state_dict(model_state['model_state_dict'])
|
| 56 |
else:
|
|
|
|
| 51 |
output_dim=args.num_classes)
|
| 52 |
|
| 53 |
model_state = torch.load(args.model_path)
|
| 54 |
+
|
| 55 |
+
if 'vocab_size' in model_state:
|
| 56 |
+
vocab_size = model_state['vocab_size']
|
| 57 |
+
tokenizer.vocab_size = vocab_size
|
| 58 |
+
|
| 59 |
if 'model_state_dict' in model_state:
|
| 60 |
model.load_state_dict(model_state['model_state_dict'])
|
| 61 |
else:
|
knowledge_distillation.py
CHANGED
|
@@ -164,7 +164,10 @@ class DistillationTrainer:
|
|
| 164 |
if val_f1 > self.best_val_f1:
|
| 165 |
self.best_val_f1 = val_f1
|
| 166 |
self.best_model_state = self.student_model.state_dict().copy()
|
| 167 |
-
torch.save(
|
|
|
|
|
|
|
|
|
|
| 168 |
logger.info(f"New best model saved with validation F1: {val_f1:.4f}")
|
| 169 |
|
| 170 |
logger.info(f"Epoch {epoch+1}/{epochs}: "
|
|
|
|
| 164 |
if val_f1 > self.best_val_f1:
|
| 165 |
self.best_val_f1 = val_f1
|
| 166 |
self.best_model_state = self.student_model.state_dict().copy()
|
| 167 |
+
torch.save({
|
| 168 |
+
'model_state_dict': self.student_model.state_dict(),
|
| 169 |
+
'vocab_size': self.student_model.vocab_size,
|
| 170 |
+
}, save_path)
|
| 171 |
logger.info(f"New best model saved with validation F1: {val_f1:.4f}")
|
| 172 |
|
| 173 |
logger.info(f"Epoch {epoch+1}/{epochs}: "
|