Commit ·
6cf4c1f
1
Parent(s): cb428cb
Forget num_classes
Browse files- distill_bert_to_lstm.py +2 -0
- knowledge_distillation.py +2 -0
- train.py +1 -0
- trainer.py +3 -1
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()}")
|