jesse-tong commited on
Commit
6cf4c1f
·
1 Parent(s): cb428cb

Forget num_classes

Browse files
distill_bert_to_lstm.py CHANGED
@@ -153,6 +153,8 @@ def main():
153
  temperature=args.temperature,
154
  alpha=args.alpha,
155
  lr=args.learning_rate,
 
 
156
  weight_decay=1e-5
157
  )
158
 
 
153
  temperature=args.temperature,
154
  alpha=args.alpha,
155
  lr=args.learning_rate,
156
+ num_categories=num_categories,
157
+ num_classes=args.num_classes,
158
  weight_decay=1e-5
159
  )
160
 
knowledge_distillation.py CHANGED
@@ -26,6 +26,7 @@ class DistillationTrainer:
26
  max_grad_norm=1.0,
27
  label_mapping=None,
28
  num_categories=1,
 
29
  device=None
30
  ):
31
  self.teacher_model = teacher_model
@@ -37,6 +38,7 @@ class DistillationTrainer:
37
  self.alpha = alpha
38
  self.max_grad_norm = max_grad_norm
39
  self.num_categories = num_categories
 
40
 
41
  self.device = device if device else torch.device('cuda' if torch.cuda.is_available() else 'cpu')
42
  logger.info(f"Using device: {self.device}")
 
26
  max_grad_norm=1.0,
27
  label_mapping=None,
28
  num_categories=1,
29
+ num_classes=2,
30
  device=None
31
  ):
32
  self.teacher_model = teacher_model
 
38
  self.alpha = alpha
39
  self.max_grad_norm = max_grad_norm
40
  self.num_categories = num_categories
41
+ self.num_classes = num_classes
42
 
43
  self.device = device if device else torch.device('cuda' if torch.cuda.is_available() else 'cpu')
44
  logger.info(f"Using device: {self.device}")
train.py CHANGED
@@ -121,6 +121,7 @@ def main():
121
  warmup_proportion=args.warmup_proportion,
122
  gradient_accumulation_steps=args.grad_accum_steps,
123
  num_categories=num_categories,
 
124
  )
125
 
126
  # Train the model
 
121
  warmup_proportion=args.warmup_proportion,
122
  gradient_accumulation_steps=args.grad_accum_steps,
123
  num_categories=num_categories,
124
+ num_classes=args.num_classes,
125
  )
126
 
127
  # Train the model
trainer.py CHANGED
@@ -28,6 +28,7 @@ class Trainer:
28
  warmup_proportion=0.1,
29
  gradient_accumulation_steps=1,
30
  max_grad_norm=1.0,
 
31
  num_categories=1,
32
  device=None
33
  ):
@@ -70,6 +71,7 @@ class Trainer:
70
  self.best_val_f1 = 0.0
71
  self.best_model_state = None
72
 
 
73
  # For training if using multiple categories (e.g., multiple sentiment classes, there can be multiple sentiment in one document)
74
  self.num_categories = num_categories
75
 
@@ -237,7 +239,7 @@ class Trainer:
237
  end_idx = (i + 1) * self.num_classes
238
  category_outputs = outputs[:, start_idx:end_idx] # Shape (batch, num_classes)
239
  category_labels = labels[:, i] # Shape (batch)
240
-
241
  # Ensure category_labels are in [0, self.num_classes - 1]
242
  if category_labels.max() >= self.num_classes or category_labels.min() < 0:
243
  print(f"ERROR: Category {i} labels out of range [0, {self.num_classes - 1}]: min={category_labels.min()}, max={category_labels.max()}")
 
28
  warmup_proportion=0.1,
29
  gradient_accumulation_steps=1,
30
  max_grad_norm=1.0,
31
+ num_classes=2,
32
  num_categories=1,
33
  device=None
34
  ):
 
71
  self.best_val_f1 = 0.0
72
  self.best_model_state = None
73
 
74
+ self.num_classes = num_classes # Number of classes for classification
75
  # For training if using multiple categories (e.g., multiple sentiment classes, there can be multiple sentiment in one document)
76
  self.num_categories = num_categories
77
 
 
239
  end_idx = (i + 1) * self.num_classes
240
  category_outputs = outputs[:, start_idx:end_idx] # Shape (batch, num_classes)
241
  category_labels = labels[:, i] # Shape (batch)
242
+
243
  # Ensure category_labels are in [0, self.num_classes - 1]
244
  if category_labels.max() >= self.num_classes or category_labels.min() < 0:
245
  print(f"ERROR: Category {i} labels out of range [0, {self.num_classes - 1}]: min={category_labels.min()}, max={category_labels.max()}")