Commit ·
80b78df
1
Parent(s): 8e3d6fe
Fix a bug in evaluation in multi-class
Browse files- knowledge_distillation.py +19 -8
- trainer.py +15 -2
knowledge_distillation.py
CHANGED
|
@@ -5,6 +5,7 @@ import numpy as np
|
|
| 5 |
from tqdm import tqdm
|
| 6 |
import logging
|
| 7 |
import os
|
|
|
|
| 8 |
|
| 9 |
logger = logging.getLogger(__name__)
|
| 10 |
|
|
@@ -189,10 +190,14 @@ class DistillationTrainer:
|
|
| 189 |
|
| 190 |
# Calculate training metrics
|
| 191 |
train_loss = train_loss / len(self.train_loader)
|
| 192 |
-
|
|
|
|
|
|
|
|
|
|
| 193 |
|
|
|
|
| 194 |
# Evaluate on validation set
|
| 195 |
-
val_loss, val_acc, val_f1 = self.evaluate()
|
| 196 |
|
| 197 |
# Update learning rate based on validation performance
|
| 198 |
self.scheduler.step(val_f1)
|
|
@@ -205,11 +210,11 @@ class DistillationTrainer:
|
|
| 205 |
'model_state_dict': self.student_model.state_dict(),
|
| 206 |
'label_mapping': self.label_mapping,
|
| 207 |
}, save_path)
|
| 208 |
-
logger.info(f"New best model saved with validation F1: {val_f1:.4f}")
|
| 209 |
|
| 210 |
logger.info(f"Epoch {epoch+1}/{epochs}: "
|
| 211 |
f"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, "
|
| 212 |
-
f"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}, Val F1: {val_f1:.4f}")
|
| 213 |
|
| 214 |
# Load best model for final evaluation
|
| 215 |
if self.best_model_state is not None:
|
|
@@ -285,11 +290,17 @@ class DistillationTrainer:
|
|
| 285 |
# Calculate metrics
|
| 286 |
eval_loss = eval_loss / len(data_loader)
|
| 287 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 288 |
# Accuracy
|
| 289 |
-
accuracy =
|
| 290 |
-
|
|
|
|
|
|
|
|
|
|
| 291 |
# F1 score (macro-averaged)
|
| 292 |
-
from sklearn.metrics import f1_score
|
| 293 |
f1 = f1_score(all_labels, all_preds, average='macro')
|
| 294 |
|
| 295 |
-
return eval_loss, accuracy, f1
|
|
|
|
| 5 |
from tqdm import tqdm
|
| 6 |
import logging
|
| 7 |
import os
|
| 8 |
+
from sklearn.metrics import f1_score, accuracy_score, precision_score, recall_score
|
| 9 |
|
| 10 |
logger = logging.getLogger(__name__)
|
| 11 |
|
|
|
|
| 190 |
|
| 191 |
# Calculate training metrics
|
| 192 |
train_loss = train_loss / len(self.train_loader)
|
| 193 |
+
if self.num_categories > 1:
|
| 194 |
+
all_labels = np.concatenate(all_labels, axis=0)
|
| 195 |
+
all_preds = np.concatenate(all_preds, axis=0)
|
| 196 |
+
#train_acc = sum(1 for p, l in zip(all_preds, all_labels) if p == l) / len(all_preds)
|
| 197 |
|
| 198 |
+
train_acc = accuracy_score(all_labels, all_preds)
|
| 199 |
# Evaluate on validation set
|
| 200 |
+
val_loss, val_acc, val_precision, val_recall, val_f1 = self.evaluate()
|
| 201 |
|
| 202 |
# Update learning rate based on validation performance
|
| 203 |
self.scheduler.step(val_f1)
|
|
|
|
| 210 |
'model_state_dict': self.student_model.state_dict(),
|
| 211 |
'label_mapping': self.label_mapping,
|
| 212 |
}, save_path)
|
| 213 |
+
logger.info(f"New best model saved with validation F1: {val_f1:.4f}, accuracy: {val_acc:.4f}")
|
| 214 |
|
| 215 |
logger.info(f"Epoch {epoch+1}/{epochs}: "
|
| 216 |
f"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, "
|
| 217 |
+
f"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}, Val Precision: {val_precision:.4f}, Val Recall: {val_recall:.4f}, Val F1: {val_f1:.4f}")
|
| 218 |
|
| 219 |
# Load best model for final evaluation
|
| 220 |
if self.best_model_state is not None:
|
|
|
|
| 290 |
# Calculate metrics
|
| 291 |
eval_loss = eval_loss / len(data_loader)
|
| 292 |
|
| 293 |
+
if self.num_categories > 1:
|
| 294 |
+
# Concatenate all labels and predictions
|
| 295 |
+
all_labels = np.concatenate(all_labels, axis=0)
|
| 296 |
+
all_preds = np.concatenate(all_preds, axis=0)
|
| 297 |
# Accuracy
|
| 298 |
+
accuracy = accuracy_score(all_labels, all_preds)
|
| 299 |
+
# Precision
|
| 300 |
+
precision = precision_score(all_labels, all_preds, average='macro')
|
| 301 |
+
# Recall
|
| 302 |
+
recall = recall_score(all_labels, all_preds, average='macro')
|
| 303 |
# F1 score (macro-averaged)
|
|
|
|
| 304 |
f1 = f1_score(all_labels, all_preds, average='macro')
|
| 305 |
|
| 306 |
+
return eval_loss, accuracy, precision, recall, f1
|
trainer.py
CHANGED
|
@@ -165,8 +165,16 @@ class Trainer:
|
|
| 165 |
|
| 166 |
# Calculate training metrics
|
| 167 |
train_loss /= len(self.train_loader)
|
| 168 |
-
|
| 169 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 170 |
|
| 171 |
# Validation phase
|
| 172 |
val_loss, val_acc, val_f1, val_precision, val_recall = self.evaluate(self.val_loader, "Validation")
|
|
@@ -271,6 +279,11 @@ class Trainer:
|
|
| 271 |
all_predictions.extend(preds.cpu().tolist())
|
| 272 |
all_labels.extend(labels.cpu().tolist())
|
| 273 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 274 |
# Calculate metrics
|
| 275 |
eval_loss /= len(data_loader)
|
| 276 |
accuracy = accuracy_score(all_labels, all_predictions)
|
|
|
|
| 165 |
|
| 166 |
# Calculate training metrics
|
| 167 |
train_loss /= len(self.train_loader)
|
| 168 |
+
if self.num_categories > 1:
|
| 169 |
+
# Flatten the list of predictions and labels
|
| 170 |
+
all_predictions = np.concatenate(all_predictions)
|
| 171 |
+
all_labels = np.concatenate(all_labels)
|
| 172 |
+
|
| 173 |
+
train_acc = accuracy_score(all_labels, all_predictions)
|
| 174 |
+
train_f1 = f1_score(all_labels, all_predictions, average='macro')
|
| 175 |
+
else:
|
| 176 |
+
train_acc = accuracy_score(all_labels, all_predictions)
|
| 177 |
+
train_f1 = f1_score(all_labels, all_predictions, average='macro')
|
| 178 |
|
| 179 |
# Validation phase
|
| 180 |
val_loss, val_acc, val_f1, val_precision, val_recall = self.evaluate(self.val_loader, "Validation")
|
|
|
|
| 279 |
all_predictions.extend(preds.cpu().tolist())
|
| 280 |
all_labels.extend(labels.cpu().tolist())
|
| 281 |
|
| 282 |
+
if self.num_categories > 1:
|
| 283 |
+
# Flatten the list of predictions and labels
|
| 284 |
+
all_predictions = np.concatenate(all_predictions)
|
| 285 |
+
all_labels = np.concatenate(all_labels)
|
| 286 |
+
|
| 287 |
# Calculate metrics
|
| 288 |
eval_loss /= len(data_loader)
|
| 289 |
accuracy = accuracy_score(all_labels, all_predictions)
|