jesse-tong commited on
Commit
80b78df
·
1 Parent(s): 8e3d6fe

Fix a bug in evaluation in multi-class

Browse files
Files changed (2) hide show
  1. knowledge_distillation.py +19 -8
  2. 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
- train_acc = sum(1 for p, l in zip(all_preds, all_labels) if p == l) / len(all_preds)
 
 
 
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 = sum(1 for p, l in zip(all_preds, all_labels) if p == l) / len(all_preds)
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
- train_acc = accuracy_score(all_labels, all_predictions)
169
- train_f1 = f1_score(all_labels, all_predictions, average='macro')
 
 
 
 
 
 
 
 
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)