jesse-tong commited on
Commit
ca56e1c
·
1 Parent(s): 01b0ca9

Save vocab size with model_state_dict

Browse files
Files changed (2) hide show
  1. inference_lstm.py +5 -0
  2. 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(self.student_model.state_dict(), save_path)
 
 
 
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}: "