bertuner 0.1.0__tar.gz → 0.1.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.1.0
3
+ Version: 0.1.2
4
4
  Summary: Hyperparameter optimization and fine-tuning for BERT-style text classifiers (Optuna + MLflow), with long-context ModernBERT support
5
5
  Author-email: elemets <alafunnell@gmail.com>
6
6
  License: MIT
@@ -24,13 +24,15 @@ Requires-Dist: torch>=2.0
24
24
  Requires-Dist: transformers>=4.48
25
25
  Requires-Dist: numpy>=1.24
26
26
  Requires-Dist: pandas>=2.0
27
+ Requires-Dist: matplotlib>=3.7
27
28
  Requires-Dist: scikit-learn>=1.3
29
+ Requires-Dist: optuna>=3.0
30
+ Requires-Dist: mlflow>=2.9
31
+ Requires-Dist: datasets>=2.14
32
+ Requires-Dist: tensorboard>=2.15
33
+ Requires-Dist: accelerate>=0.26
34
+ Requires-Dist: sentencepiece>=0.1.99
28
35
  Provides-Extra: train
29
- Requires-Dist: optuna>=3.0; extra == "train"
30
- Requires-Dist: mlflow>=2.9; extra == "train"
31
- Requires-Dist: datasets>=2.14; extra == "train"
32
- Requires-Dist: tensorboard>=2.15; extra == "train"
33
- Requires-Dist: accelerate>=0.26; extra == "train"
34
36
  Provides-Extra: dev
35
37
  Requires-Dist: pytest>=7.0; extra == "dev"
36
38
  Requires-Dist: build; extra == "dev"
@@ -41,13 +43,12 @@ Dynamic: license-file
41
43
 
42
44
  A library for hyperparameter optimization and fine-tuning of BERT-based classification models. It integrates **Optuna** for efficient search and **MLflow** for experiment tracking.
43
45
 
44
- Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens). Per-architecture dropout is applied automatically, `max_length` is clamped to each model's real context window, precision is bf16 where the GPU supports it, and gradient checkpointing switches on automatically for sequences longer than 1024 tokens (override with `gradient_checkpointing=True/False`).
46
+ Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens).
45
47
 
46
48
  ## Installation
47
49
 
48
50
  ```bash
49
- pip install bertuner[train] # training + inference
50
- pip install bertuner # inference only (BERTunePredictor)
51
+ pip install bertuner # training + inference, batteries included
51
52
  ```
52
53
 
53
54
  From source (development):
@@ -105,6 +106,18 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
105
106
 
106
107
  Grouped data (e.g. multiple notes per patient): pass `group_key="patient_id"` and the train/val/test split guarantees no group leaks across splits.
107
108
 
109
+ ### Metrics logged to MLflow
110
+
111
+ Final runs log one canonical metric set for both `Validation_*` and `Test_*`:
112
+
113
+ - **Binary:** accuracy, balanced accuracy, precision, recall, specificity, F1, Matthews correlation coefficient (MCC), average precision, AUROC, Brier score, and log loss.
114
+ - **Multiclass:** accuracy, balanced accuracy, macro and weighted precision/recall/F1, MCC, macro and weighted average precision/AUROC, and log loss.
115
+ - **Multi-label:** subset accuracy, Hamming loss/accuracy, micro/macro/sample precision/recall/F1 and Jaccard, micro MCC, micro/macro average precision and AUROC, Brier score, and log loss.
116
+
117
+ The optimized binary or per-label decision threshold is logged as a parameter. Multiclass runs log `decision_rule=argmax`. Training-time `eval_*` metrics are not duplicated in MLflow; their losses remain visible in the `plots/training_vs_evaluation_loss.png` artifact.
118
+
119
+ When optimizing a lower-is-better metric such as `log_loss`, pass `greater_is_better=False`; Optuna and best-checkpoint selection will both minimize it.
120
+
108
121
  ## Customizing the hyperparameter search
109
122
 
110
123
  Two things are configurable: **which models** are searched and **which hyperparameters** with what ranges.
@@ -2,13 +2,12 @@
2
2
 
3
3
  A library for hyperparameter optimization and fine-tuning of BERT-based classification models. It integrates **Optuna** for efficient search and **MLflow** for experiment tracking.
4
4
 
5
- Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens). Per-architecture dropout is applied automatically, `max_length` is clamped to each model's real context window, precision is bf16 where the GPU supports it, and gradient checkpointing switches on automatically for sequences longer than 1024 tokens (override with `gradient_checkpointing=True/False`).
5
+ Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens).
6
6
 
7
7
  ## Installation
8
8
 
9
9
  ```bash
10
- pip install bertuner[train] # training + inference
11
- pip install bertuner # inference only (BERTunePredictor)
10
+ pip install bertuner # training + inference, batteries included
12
11
  ```
13
12
 
14
13
  From source (development):
@@ -66,6 +65,18 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
66
65
 
67
66
  Grouped data (e.g. multiple notes per patient): pass `group_key="patient_id"` and the train/val/test split guarantees no group leaks across splits.
68
67
 
68
+ ### Metrics logged to MLflow
69
+
70
+ Final runs log one canonical metric set for both `Validation_*` and `Test_*`:
71
+
72
+ - **Binary:** accuracy, balanced accuracy, precision, recall, specificity, F1, Matthews correlation coefficient (MCC), average precision, AUROC, Brier score, and log loss.
73
+ - **Multiclass:** accuracy, balanced accuracy, macro and weighted precision/recall/F1, MCC, macro and weighted average precision/AUROC, and log loss.
74
+ - **Multi-label:** subset accuracy, Hamming loss/accuracy, micro/macro/sample precision/recall/F1 and Jaccard, micro MCC, micro/macro average precision and AUROC, Brier score, and log loss.
75
+
76
+ The optimized binary or per-label decision threshold is logged as a parameter. Multiclass runs log `decision_rule=argmax`. Training-time `eval_*` metrics are not duplicated in MLflow; their losses remain visible in the `plots/training_vs_evaluation_loss.png` artifact.
77
+
78
+ When optimizing a lower-is-better metric such as `log_loss`, pass `greater_is_better=False`; Optuna and best-checkpoint selection will both minimize it.
79
+
69
80
  ## Customizing the hyperparameter search
70
81
 
71
82
  Two things are configurable: **which models** are searched and **which hyperparameters** with what ranges.
@@ -20,12 +20,18 @@ from transformers import (
20
20
  set_seed,
21
21
  )
22
22
  from sklearn.metrics import (
23
+ accuracy_score,
24
+ average_precision_score,
25
+ balanced_accuracy_score,
26
+ brier_score_loss,
27
+ f1_score,
28
+ hamming_loss,
29
+ jaccard_score,
30
+ log_loss,
31
+ matthews_corrcoef,
23
32
  precision_score,
24
33
  recall_score,
25
- f1_score,
26
- average_precision_score,
27
34
  roc_auc_score,
28
- accuracy_score,
29
35
  )
30
36
  from sklearn.preprocessing import label_binarize
31
37
  from mlflow.tracking import MlflowClient
@@ -100,6 +106,9 @@ class BERTuneClassifier:
100
106
  # Accepts a plain path (converted to file: URI) or any mlflow URI.
101
107
  if "://" not in mlflow_tracking_uri:
102
108
  mlflow_tracking_uri = f"file:{os.path.abspath(mlflow_tracking_uri)}"
109
+ if mlflow_tracking_uri.startswith("file:"):
110
+ # mlflow >=3.14 rejects the filesystem backend unless opted in
111
+ os.environ.setdefault("MLFLOW_ALLOW_FILE_STORE", "true")
103
112
  self.mlflow_uri = mlflow_tracking_uri
104
113
  else:
105
114
  self.mlflow_uri = f"http://127.0.0.1:{mlflow_port}"
@@ -210,8 +219,12 @@ class BERTuneClassifier:
210
219
  """
211
220
  Computes metrics for both single-label and multi-label modes.
212
221
 
213
- Single-label: accuracy, f1, precision, recall, specificity, AP, AUC-ROC
214
- Multi-label: micro/macro/sample-averaged variants of the above
222
+ Binary: accuracy, balanced accuracy, precision/recall/specificity, F1,
223
+ MCC, average precision, AUROC, Brier score, and log loss.
224
+ Multiclass: macro and weighted classification metrics plus MCC,
225
+ one-vs-rest ranking metrics, and log loss.
226
+ Multi-label: micro/macro/sample classification metrics, Hamming and
227
+ Jaccard scores, ranking metrics, MCC, Brier score, and log loss.
215
228
  """
216
229
  predictions, labels = eval_pred
217
230
 
@@ -223,29 +236,46 @@ class BERTuneClassifier:
223
236
 
224
237
  metrics = {
225
238
  "accuracy": accuracy_score(labels, preds),
239
+ "hamming_loss": hamming_loss(labels, preds),
240
+ "hamming_accuracy": 1.0 - hamming_loss(labels, preds),
226
241
  "precision_micro": precision_score(
227
242
  labels, preds, average="micro", zero_division=0
228
243
  ),
244
+ "precision_macro": precision_score(
245
+ labels, preds, average="macro", zero_division=0
246
+ ),
247
+ "precision_samples": precision_score(
248
+ labels, preds, average="samples", zero_division=0
249
+ ),
229
250
  "recall_micro": recall_score(labels, preds, average="micro", zero_division=0),
251
+ "recall_macro": recall_score(labels, preds, average="macro", zero_division=0),
252
+ "recall_samples": recall_score(
253
+ labels, preds, average="samples", zero_division=0
254
+ ),
230
255
  "f1_micro": f1_score(labels, preds, average="micro", zero_division=0),
231
256
  "f1_macro": f1_score(labels, preds, average="macro", zero_division=0),
232
257
  "f1_samples": f1_score(labels, preds, average="samples", zero_division=0),
258
+ "jaccard_micro": jaccard_score(
259
+ labels, preds, average="micro", zero_division=0
260
+ ),
261
+ "jaccard_macro": jaccard_score(
262
+ labels, preds, average="macro", zero_division=0
263
+ ),
264
+ "jaccard_samples": jaccard_score(
265
+ labels, preds, average="samples", zero_division=0
266
+ ),
267
+ "mcc_micro": matthews_corrcoef(labels.ravel(), preds.ravel()),
268
+ "brier_score_micro": brier_score_loss(labels.ravel(), probs.ravel()),
269
+ "log_loss_micro": log_loss(
270
+ labels.ravel(), probs.ravel(), labels=[0, 1]
271
+ ),
233
272
  }
234
- # Average-precision and AUC per label, then macro-average
235
-
236
- aps = []
237
- aucs = []
238
-
239
- for i in range(labels.shape[1]):
240
- y_i = labels[:, i]
241
- p_i = probs[:, i]
242
-
243
- if len(np.unique(y_i)) > 1:
244
- aps.append(average_precision_score(y_i, p_i))
245
- aucs.append(roc_auc_score(y_i, p_i))
246
-
247
- metrics["avg_precision"] = float(np.mean(aps)) if aps else 0.5
248
- metrics["auc_roc"] = float(np.mean(aucs)) if aucs else 0.5
273
+ ranking = self._multilabel_ranking_metrics(labels, probs)
274
+ # Preserve the original names for optimization compatibility.
275
+ metrics["avg_precision"] = ranking["average_precision_macro"]
276
+ metrics["avg_precision_micro"] = ranking["average_precision_micro"]
277
+ metrics["auc_roc"] = ranking["auroc_macro"]
278
+ metrics["auc_roc_micro"] = ranking["auroc_micro"]
249
279
  else:
250
280
  # Single-label: predictions shape (N, num_classes) or (N,)
251
281
  if len(predictions.shape) > 1:
@@ -258,23 +288,47 @@ class BERTuneClassifier:
258
288
  if self.is_binary:
259
289
  metrics = {
260
290
  "accuracy": accuracy_score(labels, preds),
291
+ "balanced_accuracy": balanced_accuracy_score(labels, preds),
261
292
  "precision": precision_score(labels, preds, zero_division=0),
262
293
  "recall": recall_score(labels, preds, zero_division=0),
263
294
  "f1": f1_score(labels, preds, zero_division=0),
264
295
  "specificity": recall_score(labels, preds, pos_label=0, zero_division=0),
296
+ "mcc": matthews_corrcoef(labels, preds),
265
297
  }
266
298
  if probs is not None:
267
- metrics["avg_precision"] = average_precision_score(labels, probs[:, 1])
299
+ positive_probs = probs[:, 1]
300
+ metrics["avg_precision"] = average_precision_score(
301
+ labels, positive_probs
302
+ )
268
303
  metrics["auc_roc"] = (
269
- roc_auc_score(labels, probs[:, 1]) if len(np.unique(labels)) > 1 else 0.5
304
+ roc_auc_score(labels, positive_probs)
305
+ if len(np.unique(labels)) > 1
306
+ else 0.5
307
+ )
308
+ metrics["brier_score"] = brier_score_loss(
309
+ labels, positive_probs
310
+ )
311
+ metrics["log_loss"] = log_loss(
312
+ labels, positive_probs, labels=[0, 1]
270
313
  )
271
314
  else:
272
315
  # Multiclass: macro-averaged metrics, one-vs-rest ranking metrics
273
316
  metrics = {
274
317
  "accuracy": accuracy_score(labels, preds),
318
+ "balanced_accuracy": balanced_accuracy_score(labels, preds),
275
319
  "precision": precision_score(labels, preds, average="macro", zero_division=0),
320
+ "precision_weighted": precision_score(
321
+ labels, preds, average="weighted", zero_division=0
322
+ ),
276
323
  "recall": recall_score(labels, preds, average="macro", zero_division=0),
324
+ "recall_weighted": recall_score(
325
+ labels, preds, average="weighted", zero_division=0
326
+ ),
277
327
  "f1": f1_score(labels, preds, average="macro", zero_division=0),
328
+ "f1_weighted": f1_score(
329
+ labels, preds, average="weighted", zero_division=0
330
+ ),
331
+ "mcc": matthews_corrcoef(labels, preds),
278
332
  }
279
333
  # Ranking metrics need every class present in the eval split
280
334
  if probs is not None and len(np.unique(labels)) == probs.shape[1]:
@@ -282,15 +336,67 @@ class BERTuneClassifier:
282
336
  metrics["avg_precision"] = average_precision_score(
283
337
  y_bin, probs, average="macro"
284
338
  )
339
+ metrics["avg_precision_weighted"] = average_precision_score(
340
+ y_bin, probs, average="weighted"
341
+ )
285
342
  metrics["auc_roc"] = roc_auc_score(
286
343
  labels, probs, multi_class="ovr", average="macro"
287
344
  )
345
+ metrics["auc_roc_weighted"] = roc_auc_score(
346
+ labels, probs, multi_class="ovr", average="weighted"
347
+ )
288
348
  elif probs is not None:
289
349
  metrics["avg_precision"] = 0.5
350
+ metrics["avg_precision_weighted"] = 0.5
290
351
  metrics["auc_roc"] = 0.5
352
+ metrics["auc_roc_weighted"] = 0.5
353
+ if probs is not None:
354
+ metrics["log_loss"] = log_loss(
355
+ labels,
356
+ probs,
357
+ labels=np.arange(probs.shape[1]),
358
+ )
291
359
 
292
360
  return metrics
293
361
 
362
+ @staticmethod
363
+ def _multilabel_ranking_metrics(labels, probs):
364
+ """Returns robust macro/micro AP and AUROC for multi-label targets."""
365
+ valid_columns = [
366
+ i for i in range(labels.shape[1]) if len(np.unique(labels[:, i])) > 1
367
+ ]
368
+ if valid_columns:
369
+ aps = [
370
+ average_precision_score(labels[:, i], probs[:, i])
371
+ for i in valid_columns
372
+ ]
373
+ aucs = [
374
+ roc_auc_score(labels[:, i], probs[:, i]) for i in valid_columns
375
+ ]
376
+ average_precision_macro = float(np.mean(aps))
377
+ auroc_macro = float(np.mean(aucs))
378
+ else:
379
+ average_precision_macro = 0.5
380
+ auroc_macro = 0.5
381
+
382
+ flat_labels = labels.ravel()
383
+ flat_probs = probs.ravel()
384
+ if len(np.unique(flat_labels)) > 1:
385
+ average_precision_micro = average_precision_score(
386
+ flat_labels, flat_probs
387
+ )
388
+ auroc_micro = roc_auc_score(flat_labels, flat_probs)
389
+ else:
390
+ average_precision_micro = 0.5
391
+ auroc_micro = 0.5
392
+
393
+ return {
394
+ "average_precision_macro": average_precision_macro,
395
+ "average_precision_micro": float(average_precision_micro),
396
+ "auroc_macro": auroc_macro,
397
+ "auroc_micro": float(auroc_micro),
398
+ }
399
+
294
400
  # ------------------------------------------------------------------
295
401
  # Data preparation
296
402
  # ------------------------------------------------------------------
@@ -496,6 +602,7 @@ class BERTuneClassifier:
496
602
  num_train_epochs=6,
497
603
  save_total_limit=2,
498
604
  load_best_model_at_end=True,
605
+ logging_strategy="epoch",
499
606
  seed=self.seed,
500
607
  remove_unused_columns=True,
501
608
  report_to=["none"],
@@ -522,6 +629,7 @@ class BERTuneClassifier:
522
629
  trainer.train()
523
630
  metrics = trainer.evaluate()
524
631
  mlflow.log_metrics(metrics)
632
+ self._log_loss_curve(trainer.state.log_history)
525
633
  else:
526
634
  trainer.train()
527
635
  metrics = trainer.evaluate()
@@ -556,7 +664,12 @@ class BERTuneClassifier:
556
664
  mlflow.set_experiment(experiment_name=exp_name)
557
665
 
558
666
  sampler = TPESampler(seed=self.seed)
559
- study = optuna.create_study(direction="maximize", study_name=study_name, sampler=sampler)
667
+ direction = "maximize" if greater_is_better else "minimize"
668
+ study = optuna.create_study(
669
+ direction=direction,
670
+ study_name=study_name,
671
+ sampler=sampler,
672
+ )
560
673
  self.optimize_metric = optimize_metric
561
674
  study.optimize(self._objective, n_trials=n_trials)
562
675
 
@@ -603,7 +716,12 @@ class BERTuneClassifier:
603
716
  greater_is_better=self.greater_is_better,
604
717
  save_total_limit=2,
605
718
  load_best_model_at_end=True,
606
- report_to=["tensorboard", "mlflow"],
719
+ logging_strategy="epoch",
720
+ # The final validation/test metrics are logged explicitly below.
721
+ # Sending Trainer logs to MLflow as well creates a second, misleading
722
+ # eval_* metric set from the last epoch rather than the restored best
723
+ # checkpoint.
724
+ report_to=["tensorboard"],
607
725
  logging_dir=f"{final_dir}/logs",
608
726
  seed=self.seed,
609
727
  **self._precision_flags(),
@@ -642,14 +760,70 @@ class BERTuneClassifier:
642
760
  )
643
761
 
644
762
  mlflow.log_params(self.best_params)
645
- for _, row in metrics_df.iterrows():
646
- for m in ["Accuracy", "F1", "AUC"]:
647
- mlflow.log_metric(f"{row['Split']}_{m}", row[m])
763
+ self._log_final_metrics(metrics_df)
764
+ self._log_loss_curve(trainer.state.log_history)
648
765
 
649
766
  self._save_model(final_dir, trainer, tokenizer, model_path, max_length)
650
767
 
651
768
  return metrics_df, model, test_ds
652
769
 
770
+ def _log_final_metrics(self, metrics_df):
771
+ """Logs every numeric validation/test result as a canonical MLflow metric."""
772
+ for _, row in metrics_df.iterrows():
773
+ for column, value in row.items():
774
+ if column in {"Split", "Threshold"}:
775
+ continue
776
+ if isinstance(value, (int, float, np.integer, np.floating)) and np.isfinite(
777
+ value
778
+ ):
779
+ mlflow.log_metric(f"{row['Split']}_{column}", float(value))
780
+
781
+ threshold = metrics_df.iloc[0].get("Threshold")
782
+ if threshold is None or (
783
+ isinstance(threshold, (float, np.floating)) and np.isnan(threshold)
784
+ ):
785
+ mlflow.log_param("decision_rule", "argmax")
786
+ else:
787
+ mlflow.log_param("decision_threshold", threshold)
788
+
789
+ def _log_loss_curve(self, log_history):
790
+ """Logs an epoch-level training-vs-evaluation loss plot to MLflow."""
791
+ train_points = [
792
+ (entry.get("epoch", entry.get("step")), entry["loss"])
793
+ for entry in log_history
794
+ if "loss" in entry
795
+ ]
796
+ eval_points = [
797
+ (entry.get("epoch", entry.get("step")), entry["eval_loss"])
798
+ for entry in log_history
799
+ if "eval_loss" in entry
800
+ ]
801
+
802
+ if not train_points or not eval_points:
803
+ return
804
+
805
+ # Import lazily so metric-only use of BERTuner does not initialise a
806
+ # plotting backend or matplotlib cache.
807
+ import matplotlib.pyplot as plt
808
+
809
+ fig, ax = plt.subplots(figsize=(8, 5))
810
+ try:
811
+ train_x, train_loss = zip(*train_points)
812
+ eval_x, eval_loss = zip(*eval_points)
813
+ ax.plot(train_x, train_loss, marker="o", label="Training loss")
814
+ ax.plot(eval_x, eval_loss, marker="o", label="Evaluation loss")
815
+ ax.set(
816
+ title="Training vs Evaluation Loss",
817
+ xlabel="Epoch",
818
+ ylabel="Loss",
819
+ )
820
+ ax.grid(alpha=0.3)
821
+ ax.legend()
822
+ fig.tight_layout()
823
+ mlflow.log_figure(fig, "plots/training_vs_evaluation_loss.png")
824
+ finally:
825
+ plt.close(fig)
826
+
653
827
  # ------------------------------------------------------------------
654
828
  # Threshold optimisation
655
829
  # ------------------------------------------------------------------
@@ -719,18 +893,42 @@ class BERTuneClassifier:
719
893
  if self.is_multilabel:
720
894
  # thresh is shape (num_labels,); p is (N, num_labels)
721
895
  preds = (p >= thresh).astype(int)
896
+ ranking = self._multilabel_ranking_metrics(y, p)
897
+ current_hamming_loss = hamming_loss(y, preds)
722
898
  return {
723
899
  "Split": split,
724
900
  "Accuracy": accuracy_score(y, preds),
901
+ "Hamming_Loss": current_hamming_loss,
902
+ "Hamming_Accuracy": 1.0 - current_hamming_loss,
725
903
  "Precision_micro": precision_score(y, preds, average="micro", zero_division=0),
904
+ "Precision_macro": precision_score(y, preds, average="macro", zero_division=0),
905
+ "Precision_samples": precision_score(
906
+ y, preds, average="samples", zero_division=0
907
+ ),
726
908
  "Recall_micro": recall_score(y, preds, average="micro", zero_division=0),
909
+ "Recall_macro": recall_score(y, preds, average="macro", zero_division=0),
910
+ "Recall_samples": recall_score(
911
+ y, preds, average="samples", zero_division=0
912
+ ),
727
913
  "F1": f1_score(y, preds, average="macro", zero_division=0),
728
914
  "F1_micro": f1_score(y, preds, average="micro", zero_division=0),
729
915
  "F1_samples": f1_score(y, preds, average="samples", zero_division=0),
730
- "AP": average_precision_score(y, p, average="macro"),
731
- "AUC": (
732
- roc_auc_score(y, p, average="macro") if len(np.unique(y)) > 1 else 0.5
916
+ "Jaccard_micro": jaccard_score(
917
+ y, preds, average="micro", zero_division=0
918
+ ),
919
+ "Jaccard_macro": jaccard_score(
920
+ y, preds, average="macro", zero_division=0
921
+ ),
922
+ "Jaccard_samples": jaccard_score(
923
+ y, preds, average="samples", zero_division=0
733
924
  ),
925
+ "MCC_micro": matthews_corrcoef(y.ravel(), preds.ravel()),
926
+ "Average_Precision_macro": ranking["average_precision_macro"],
927
+ "Average_Precision_micro": ranking["average_precision_micro"],
928
+ "AUROC_macro": ranking["auroc_macro"],
929
+ "AUROC_micro": ranking["auroc_micro"],
930
+ "Brier_Score_micro": brier_score_loss(y.ravel(), p.ravel()),
931
+ "Log_Loss_micro": log_loss(y.ravel(), p.ravel(), labels=[0, 1]),
734
932
  "Threshold": str(np.round(thresh, 3).tolist()),
735
933
  }
736
934
  elif self.is_binary:
@@ -739,12 +937,16 @@ class BERTuneClassifier:
739
937
  return {
740
938
  "Split": split,
741
939
  "Accuracy": accuracy_score(y, preds),
940
+ "Balanced_Accuracy": balanced_accuracy_score(y, preds),
742
941
  "Precision": precision_score(y, preds, zero_division=0),
743
942
  "Recall": recall_score(y, preds, zero_division=0),
744
943
  "F1": f1_score(y, preds, zero_division=0),
745
944
  "Specificity": recall_score(y, preds, pos_label=0, zero_division=0),
746
- "AP": average_precision_score(y, p),
747
- "AUC": roc_auc_score(y, p) if len(np.unique(y)) > 1 else 0.5,
945
+ "MCC": matthews_corrcoef(y, preds),
946
+ "Average_Precision": average_precision_score(y, p),
947
+ "AUROC": roc_auc_score(y, p) if len(np.unique(y)) > 1 else 0.5,
948
+ "Brier_Score": brier_score_loss(y, p),
949
+ "Log_Loss": log_loss(y, p, labels=[0, 1]),
748
950
  "Threshold": thresh,
749
951
  }
750
952
  else:
@@ -754,10 +956,21 @@ class BERTuneClassifier:
754
956
  return {
755
957
  "Split": split,
756
958
  "Accuracy": accuracy_score(y, preds),
959
+ "Balanced_Accuracy": balanced_accuracy_score(y, preds),
757
960
  "Precision": precision_score(y, preds, average="macro", zero_division=0),
961
+ "Precision_weighted": precision_score(
962
+ y, preds, average="weighted", zero_division=0
963
+ ),
758
964
  "Recall": recall_score(y, preds, average="macro", zero_division=0),
965
+ "Recall_weighted": recall_score(
966
+ y, preds, average="weighted", zero_division=0
967
+ ),
759
968
  "F1": f1_score(y, preds, average="macro", zero_division=0),
760
- "AP": (
969
+ "F1_weighted": f1_score(
970
+ y, preds, average="weighted", zero_division=0
971
+ ),
972
+ "MCC": matthews_corrcoef(y, preds),
973
+ "Average_Precision_macro": (
761
974
  average_precision_score(
762
975
  label_binarize(y, classes=np.arange(p.shape[1])),
763
976
  p,
@@ -766,11 +979,28 @@ class BERTuneClassifier:
766
979
  if all_classes_present
767
980
  else 0.5
768
981
  ),
769
- "AUC": (
982
+ "Average_Precision_weighted": (
983
+ average_precision_score(
984
+ label_binarize(y, classes=np.arange(p.shape[1])),
985
+ p,
986
+ average="weighted",
987
+ )
988
+ if all_classes_present
989
+ else 0.5
990
+ ),
991
+ "AUROC_macro": (
770
992
  roc_auc_score(y, p, multi_class="ovr", average="macro")
771
993
  if all_classes_present
772
994
  else 0.5
773
995
  ),
996
+ "AUROC_weighted": (
997
+ roc_auc_score(y, p, multi_class="ovr", average="weighted")
998
+ if all_classes_present
999
+ else 0.5
1000
+ ),
1001
+ "Log_Loss": log_loss(
1002
+ y, p, labels=np.arange(p.shape[1])
1003
+ ),
774
1004
  "Threshold": None,
775
1005
  }
776
1006
 
@@ -0,0 +1,8 @@
1
+ """bertuner: hyperparameter optimization and fine-tuning for BERT-style text classifiers."""
2
+
3
+ __version__ = "0.1.2"
4
+
5
+ from bertuner.BERTuner import BERTuneClassifier
6
+ from bertuner.Predictor import BERTunePredictor
7
+
8
+ __all__ = ["BERTuneClassifier", "BERTunePredictor", "__version__"]
@@ -6,8 +6,6 @@ DEFAULT_MODEL_CHOICES = {
6
6
  "bert-base": "bert-base-uncased",
7
7
  "electra-small": "google/electra-small-discriminator",
8
8
  "electra-base": "google/electra-base-discriminator",
9
- "modernbert-base": "answerdotai/ModernBERT-base",
10
- "modernbert-large": "answerdotai/ModernBERT-large",
11
9
  }
12
10
 
13
11
  # Dropout attribute names per architecture (config.model_type).
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.1.0
3
+ Version: 0.1.2
4
4
  Summary: Hyperparameter optimization and fine-tuning for BERT-style text classifiers (Optuna + MLflow), with long-context ModernBERT support
5
5
  Author-email: elemets <alafunnell@gmail.com>
6
6
  License: MIT
@@ -24,13 +24,15 @@ Requires-Dist: torch>=2.0
24
24
  Requires-Dist: transformers>=4.48
25
25
  Requires-Dist: numpy>=1.24
26
26
  Requires-Dist: pandas>=2.0
27
+ Requires-Dist: matplotlib>=3.7
27
28
  Requires-Dist: scikit-learn>=1.3
29
+ Requires-Dist: optuna>=3.0
30
+ Requires-Dist: mlflow>=2.9
31
+ Requires-Dist: datasets>=2.14
32
+ Requires-Dist: tensorboard>=2.15
33
+ Requires-Dist: accelerate>=0.26
34
+ Requires-Dist: sentencepiece>=0.1.99
28
35
  Provides-Extra: train
29
- Requires-Dist: optuna>=3.0; extra == "train"
30
- Requires-Dist: mlflow>=2.9; extra == "train"
31
- Requires-Dist: datasets>=2.14; extra == "train"
32
- Requires-Dist: tensorboard>=2.15; extra == "train"
33
- Requires-Dist: accelerate>=0.26; extra == "train"
34
36
  Provides-Extra: dev
35
37
  Requires-Dist: pytest>=7.0; extra == "dev"
36
38
  Requires-Dist: build; extra == "dev"
@@ -41,13 +43,12 @@ Dynamic: license-file
41
43
 
42
44
  A library for hyperparameter optimization and fine-tuning of BERT-based classification models. It integrates **Optuna** for efficient search and **MLflow** for experiment tracking.
43
45
 
44
- Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens). Per-architecture dropout is applied automatically, `max_length` is clamped to each model's real context window, precision is bf16 where the GPU supports it, and gradient checkpointing switches on automatically for sequences longer than 1024 tokens (override with `gradient_checkpointing=True/False`).
46
+ Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as **ModernBERT** (8192 tokens).
45
47
 
46
48
  ## Installation
47
49
 
48
50
  ```bash
49
- pip install bertuner[train] # training + inference
50
- pip install bertuner # inference only (BERTunePredictor)
51
+ pip install bertuner # training + inference, batteries included
51
52
  ```
52
53
 
53
54
  From source (development):
@@ -105,6 +106,18 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
105
106
 
106
107
  Grouped data (e.g. multiple notes per patient): pass `group_key="patient_id"` and the train/val/test split guarantees no group leaks across splits.
107
108
 
109
+ ### Metrics logged to MLflow
110
+
111
+ Final runs log one canonical metric set for both `Validation_*` and `Test_*`:
112
+
113
+ - **Binary:** accuracy, balanced accuracy, precision, recall, specificity, F1, Matthews correlation coefficient (MCC), average precision, AUROC, Brier score, and log loss.
114
+ - **Multiclass:** accuracy, balanced accuracy, macro and weighted precision/recall/F1, MCC, macro and weighted average precision/AUROC, and log loss.
115
+ - **Multi-label:** subset accuracy, Hamming loss/accuracy, micro/macro/sample precision/recall/F1 and Jaccard, micro MCC, micro/macro average precision and AUROC, Brier score, and log loss.
116
+
117
+ The optimized binary or per-label decision threshold is logged as a parameter. Multiclass runs log `decision_rule=argmax`. Training-time `eval_*` metrics are not duplicated in MLflow; their losses remain visible in the `plots/training_vs_evaluation_loss.png` artifact.
118
+
119
+ When optimizing a lower-is-better metric such as `log_loss`, pass `greater_is_better=False`; Optuna and best-checkpoint selection will both minimize it.
120
+
108
121
  ## Customizing the hyperparameter search
109
122
 
110
123
  Two things are configurable: **which models** are searched and **which hyperparameters** with what ranges.
@@ -2,7 +2,14 @@ torch>=2.0
2
2
  transformers>=4.48
3
3
  numpy>=1.24
4
4
  pandas>=2.0
5
+ matplotlib>=3.7
5
6
  scikit-learn>=1.3
7
+ optuna>=3.0
8
+ mlflow>=2.9
9
+ datasets>=2.14
10
+ tensorboard>=2.15
11
+ accelerate>=0.26
12
+ sentencepiece>=0.1.99
6
13
 
7
14
  [dev]
8
15
  pytest>=7.0
@@ -10,8 +17,3 @@ build
10
17
  twine
11
18
 
12
19
  [train]
13
- optuna>=3.0
14
- mlflow>=2.9
15
- datasets>=2.14
16
- tensorboard>=2.15
17
- accelerate>=0.26
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "bertuner"
7
- version = "0.1.0"
7
+ version = "0.1.2"
8
8
  description = "Hyperparameter optimization and fine-tuning for BERT-style text classifiers (Optuna + MLflow), with long-context ModernBERT support"
9
9
  readme = "README.md"
10
10
  license = { text = "MIT" }
@@ -31,23 +31,27 @@ classifiers = [
31
31
  "Programming Language :: Python :: 3.12",
32
32
  "Topic :: Scientific/Engineering :: Artificial Intelligence",
33
33
  ]
34
- # Inference-only footprint; training pulls the [train] extra
35
34
  dependencies = [
36
35
  "torch>=2.0",
37
36
  "transformers>=4.48",
38
37
  "numpy>=1.24",
39
38
  "pandas>=2.0",
39
+ "matplotlib>=3.7",
40
40
  "scikit-learn>=1.3",
41
- ]
42
-
43
- [project.optional-dependencies]
44
- train = [
45
41
  "optuna>=3.0",
46
42
  "mlflow>=2.9",
47
43
  "datasets>=2.14",
48
44
  "tensorboard>=2.15",
49
45
  "accelerate>=0.26",
46
+ # transformers>=5 needs this to load tokenizers from repos that only
47
+ # ship vocab files (no tokenizer.json), e.g. Bio_ClinicalBERT
48
+ "sentencepiece>=0.1.99",
50
49
  ]
50
+
51
+ [project.optional-dependencies]
52
+ # Kept as an alias so `pip install bertuner[train]` keeps working;
53
+ # everything is already in the base dependencies.
54
+ train = []
51
55
  dev = ["pytest>=7.0", "build", "twine"]
52
56
 
53
57
  [project.urls]
@@ -1,5 +1,6 @@
1
1
  """Unit tests for BERTuneClassifier — construction, metrics, thresholds, saving."""
2
2
  import json
3
+ from pathlib import Path
3
4
  from types import SimpleNamespace
4
5
  from unittest.mock import MagicMock
5
6
 
@@ -7,6 +8,8 @@ import numpy as np
7
8
  import pandas as pd
8
9
  import pytest
9
10
  import optuna
11
+ import mlflow
12
+ from mlflow.tracking import MlflowClient
10
13
 
11
14
  from bertuner.BERTuner import BERTuneClassifier
12
15
  from bertuner.constants import (
@@ -109,8 +112,12 @@ class TestComputeMetrics:
109
112
  logits = np.array([[5.0, -5.0], [-5.0, 5.0], [5.0, -5.0], [-5.0, 5.0]])
110
113
  m = clf._compute_metrics((logits, labels))
111
114
  assert m["accuracy"] == 1.0
115
+ assert m["balanced_accuracy"] == 1.0
112
116
  assert m["f1"] == 1.0
117
+ assert m["mcc"] == 1.0
113
118
  assert m["auc_roc"] == 1.0
119
+ assert m["brier_score"] < 0.001
120
+ assert m["log_loss"] < 0.001
114
121
 
115
122
  def test_multilabel_perfect_predictions(self, tmp_path):
116
123
  cols = ["l1", "l2"]
@@ -122,6 +129,10 @@ class TestComputeMetrics:
122
129
  m = clf._compute_metrics((logits, labels))
123
130
  assert m["f1_micro"] == 1.0
124
131
  assert m["f1_macro"] == 1.0
132
+ assert m["hamming_loss"] == 0.0
133
+ assert m["jaccard_samples"] == 0.75
134
+ assert m["mcc_micro"] == 1.0
135
+ assert m["avg_precision_micro"] == 1.0
125
136
  assert m["auc_roc"] == 1.0
126
137
 
127
138
 
@@ -240,6 +251,35 @@ class TestSuggestHyperparams:
240
251
  assert isinstance(params["early_stopping_patience"], int)
241
252
  assert 3 <= params["early_stopping_patience"] <= 8
242
253
 
254
+ def test_optimize_uses_minimize_for_lower_is_better_metric(
255
+ self, tmp_path, monkeypatch
256
+ ):
257
+ clf = make_classifier(tmp_path)
258
+ client = MagicMock()
259
+ client.get_experiment_by_name.return_value = SimpleNamespace(
260
+ lifecycle_stage="active"
261
+ )
262
+ study = MagicMock(
263
+ best_params={"model": "bert-base"},
264
+ best_value=0.1,
265
+ best_trial=SimpleNamespace(number=0),
266
+ )
267
+ create_study = MagicMock(return_value=study)
268
+
269
+ monkeypatch.setattr("bertuner.BERTuner.MlflowClient", lambda: client)
270
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.set_tracking_uri", MagicMock())
271
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.set_experiment", MagicMock())
272
+ monkeypatch.setattr("bertuner.BERTuner.optuna.create_study", create_study)
273
+
274
+ clf.optimize(
275
+ n_trials=2,
276
+ optimize_metric="log_loss",
277
+ greater_is_better=False,
278
+ )
279
+
280
+ assert create_study.call_args.kwargs["direction"] == "minimize"
281
+ study.optimize.assert_called_once_with(clf._objective, n_trials=2)
282
+
243
283
 
244
284
  # ------------------------------------------------------------------
245
285
  # Metrics DataFrame
@@ -254,6 +294,11 @@ class TestBuildMetricsDf:
254
294
  df = clf._build_metrics_df(y, p, y, p, thresh=0.5)
255
295
  assert list(df["Split"]) == ["Validation", "Test"]
256
296
  assert df.iloc[0]["F1"] == 1.0
297
+ assert df.iloc[0]["Balanced_Accuracy"] == 1.0
298
+ assert df.iloc[0]["MCC"] == 1.0
299
+ assert df.iloc[0]["Average_Precision"] == 1.0
300
+ assert df.iloc[0]["AUROC"] == 1.0
301
+ assert df.iloc[0]["Brier_Score"] < 0.05
257
302
  assert df.iloc[0]["Threshold"] == 0.5
258
303
 
259
304
  def test_multilabel_uses_array_threshold(self, tmp_path):
@@ -266,6 +311,111 @@ class TestBuildMetricsDf:
266
311
  df = clf._build_metrics_df(y, p, y, p, thresh=np.array([0.5, 0.5]))
267
312
  assert df.iloc[0]["F1"] == 1.0
268
313
  assert "F1_micro" in df.columns
314
+ assert df.iloc[0]["Hamming_Loss"] == 0.0
315
+ assert df.iloc[0]["Average_Precision_macro"] == 1.0
316
+ assert df.iloc[0]["AUROC_micro"] == 1.0
317
+
318
+
319
+ class TestLossCurveLogging:
320
+ def test_logs_average_precision_and_auroc(self, tmp_path, monkeypatch):
321
+ clf = make_classifier(tmp_path)
322
+ logged = MagicMock()
323
+ logged_param = MagicMock()
324
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.log_metric", logged)
325
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.log_param", logged_param)
326
+ metrics = pd.DataFrame(
327
+ [
328
+ {
329
+ "Split": "Validation",
330
+ "Accuracy": 0.7,
331
+ "Balanced_Accuracy": 0.68,
332
+ "Precision": 0.72,
333
+ "F1": 0.6,
334
+ "Average_Precision": 0.8,
335
+ "AUROC": 0.9,
336
+ "Threshold": 0.42,
337
+ },
338
+ {
339
+ "Split": "Test",
340
+ "Accuracy": 0.65,
341
+ "Balanced_Accuracy": 0.63,
342
+ "Precision": 0.67,
343
+ "F1": 0.55,
344
+ "Average_Precision": 0.75,
345
+ "AUROC": 0.85,
346
+ "Threshold": 0.42,
347
+ },
348
+ ]
349
+ )
350
+
351
+ clf._log_final_metrics(metrics)
352
+
353
+ calls = {call.args[0]: call.args[1] for call in logged.call_args_list}
354
+ assert calls["Validation_Average_Precision"] == 0.8
355
+ assert calls["Validation_AUROC"] == 0.9
356
+ assert calls["Test_Average_Precision"] == 0.75
357
+ assert calls["Test_AUROC"] == 0.85
358
+ assert calls["Validation_Balanced_Accuracy"] == 0.68
359
+ assert calls["Test_Precision"] == 0.67
360
+ logged_param.assert_called_once_with("decision_threshold", 0.42)
361
+
362
+ def test_logs_training_and_eval_loss_figure(self, tmp_path, monkeypatch):
363
+ clf = make_classifier(tmp_path)
364
+ logged = MagicMock()
365
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.log_figure", logged)
366
+
367
+ clf._log_loss_curve(
368
+ [
369
+ {"loss": 0.8, "epoch": 1.0},
370
+ {"eval_loss": 0.9, "epoch": 1.0},
371
+ {"loss": 0.5, "epoch": 2.0},
372
+ {"eval_loss": 0.7, "epoch": 2.0},
373
+ ]
374
+ )
375
+
376
+ logged.assert_called_once()
377
+ assert logged.call_args.args[1] == "plots/training_vs_evaluation_loss.png"
378
+
379
+ def test_loss_figure_is_stored_as_mlflow_artifact(self, tmp_path):
380
+ clf = make_classifier(tmp_path)
381
+ previous_uri = mlflow.get_tracking_uri()
382
+ tracking_uri = (tmp_path / "artifact-test-mlruns").as_uri()
383
+ artifact_location = (tmp_path / "artifact-test-files").as_uri()
384
+
385
+ try:
386
+ mlflow.set_tracking_uri(tracking_uri)
387
+ experiment_id = mlflow.create_experiment(
388
+ "loss-curve-artifact-test",
389
+ artifact_location=artifact_location,
390
+ )
391
+ with mlflow.start_run(experiment_id=experiment_id) as run:
392
+ clf._log_loss_curve(
393
+ [
394
+ {"loss": 0.8, "epoch": 1.0},
395
+ {"eval_loss": 0.9, "epoch": 1.0},
396
+ ]
397
+ )
398
+
399
+ artifact_path = MlflowClient().download_artifacts(
400
+ run.info.run_id,
401
+ "plots/training_vs_evaluation_loss.png",
402
+ str(tmp_path / "downloaded-artifacts"),
403
+ )
404
+ assert Path(artifact_path).is_file()
405
+ finally:
406
+ mlflow.end_run()
407
+ mlflow.set_tracking_uri(previous_uri)
408
+
409
+ def test_skips_figure_when_either_loss_series_is_missing(
410
+ self, tmp_path, monkeypatch
411
+ ):
412
+ clf = make_classifier(tmp_path)
413
+ logged = MagicMock()
414
+ monkeypatch.setattr("bertuner.BERTuner.mlflow.log_figure", logged)
415
+
416
+ clf._log_loss_curve([{"eval_loss": 0.9, "epoch": 1.0}])
417
+
418
+ logged.assert_not_called()
269
419
 
270
420
 
271
421
  # ------------------------------------------------------------------
@@ -450,8 +600,13 @@ class TestMulticlass:
450
600
  labels = np.array([0, 1, 2, 0])
451
601
  metrics = clf._compute_metrics((logits, labels))
452
602
  assert metrics["accuracy"] == 1.0
603
+ assert metrics["balanced_accuracy"] == 1.0
453
604
  assert metrics["f1"] == 1.0
605
+ assert metrics["f1_weighted"] == 1.0
606
+ assert metrics["mcc"] == 1.0
454
607
  assert metrics["auc_roc"] == 1.0
608
+ assert metrics["auc_roc_weighted"] == 1.0
609
+ assert metrics["log_loss"] < 0.1
455
610
  assert "specificity" not in metrics
456
611
 
457
612
  def test_compute_metrics_missing_class_falls_back(self, tmp_path):
@@ -483,5 +638,105 @@ class TestMulticlass:
483
638
  )
484
639
  df = clf._build_metrics_df(y, p, y, p, thresh=None)
485
640
  assert df.iloc[0]["Accuracy"] == 1.0
641
+ assert df.iloc[0]["Balanced_Accuracy"] == 1.0
486
642
  assert df.iloc[0]["F1"] == 1.0
643
+ assert df.iloc[0]["F1_weighted"] == 1.0
644
+ assert df.iloc[0]["MCC"] == 1.0
645
+ assert df.iloc[0]["Average_Precision_macro"] == 1.0
646
+ assert df.iloc[0]["AUROC_weighted"] == 1.0
487
647
  assert df.iloc[0]["Threshold"] is None
648
+
649
+
650
+ # ------------------------------------------------------------------
651
+ # Multiclass weighted loss (regression: pre-0.1.0 code always built
652
+ # 2-element class weights, crashing cross_entropy on 3+ classes)
653
+ # ------------------------------------------------------------------
654
+
655
+
656
+ class TestMulticlassWeightedLoss:
657
+ """Weighted CE with 3 classes must not raise the shape-[2] weight error."""
658
+
659
+ def _trainer(self, class_weights, loss_type="weighted"):
660
+ import torch
661
+ from bertuner.CustomTrainer import CustomTrainer
662
+
663
+ t = CustomTrainer.__new__(CustomTrainer)
664
+ t.loss_type = loss_type
665
+ t.class_weights = class_weights
666
+ return t
667
+
668
+ def test_weighted_ce_three_classes(self, tmp_path):
669
+ import torch
670
+
671
+ clf = make_classifier(tmp_path, dataframe=make_df(num_classes=3))
672
+ ds = [{"labels": torch.tensor(i % 3)} for i in range(9)]
673
+ weights = clf._compute_class_weights(ds)
674
+ assert weights.shape == (3,)
675
+
676
+ trainer = self._trainer(weights)
677
+ logits = torch.randn(6, 3)
678
+ labels = torch.tensor([0, 1, 2, 0, 1, 2])
679
+ loss = trainer._singlelabel_loss(logits, labels, torch.device("cpu"))
680
+ assert torch.isfinite(loss)
681
+
682
+ def test_two_element_weights_reproduce_reported_error(self):
683
+ import torch
684
+
685
+ trainer = self._trainer(torch.tensor([1.0, 3.0]))
686
+ logits = torch.randn(6, 3)
687
+ labels = torch.tensor([0, 1, 2, 0, 1, 2])
688
+ with pytest.raises(RuntimeError, match="weight tensor"):
689
+ trainer._singlelabel_loss(logits, labels, torch.device("cpu"))
690
+
691
+ def test_focal_and_label_smoothing_three_classes(self, tmp_path):
692
+ import torch
693
+
694
+ logits = torch.randn(6, 3)
695
+ labels = torch.tensor([0, 1, 2, 0, 1, 2])
696
+ for loss_type in ("focal", "label_smoothing", "plain"):
697
+ trainer = self._trainer(None, loss_type=loss_type)
698
+ loss = trainer._singlelabel_loss(logits, labels, torch.device("cpu"))
699
+ assert torch.isfinite(loss)
700
+
701
+ def test_end_to_end_training_three_classes(self, tmp_path):
702
+ import torch
703
+ from transformers import (
704
+ AutoModelForSequenceClassification,
705
+ AutoTokenizer,
706
+ DataCollatorWithPadding,
707
+ TrainingArguments,
708
+ )
709
+ from bertuner.CustomTrainer import CustomTrainer
710
+
711
+ model_path = "prajjwal1/bert-tiny"
712
+ clf = make_classifier(tmp_path, dataframe=make_df(n=60, num_classes=3))
713
+ assert clf.num_labels == 3
714
+
715
+ tokenizer = AutoTokenizer.from_pretrained(model_path)
716
+ train_ds, val_ds, _ = clf._prepare_datasets(tokenizer, None, max_length=32)
717
+ class_weights = clf._compute_class_weights(train_ds)
718
+ assert class_weights.shape == (3,)
719
+
720
+ model = AutoModelForSequenceClassification.from_pretrained(
721
+ model_path, num_labels=3
722
+ )
723
+ args = TrainingArguments(
724
+ output_dir=str(tmp_path / "e2e"),
725
+ per_device_train_batch_size=8,
726
+ num_train_epochs=1,
727
+ eval_strategy="no",
728
+ save_strategy="no",
729
+ report_to=["none"],
730
+ seed=42,
731
+ )
732
+ trainer = CustomTrainer(
733
+ model=model,
734
+ args=args,
735
+ train_dataset=train_ds,
736
+ eval_dataset=val_ds,
737
+ data_collator=DataCollatorWithPadding(tokenizer),
738
+ loss_type="weighted",
739
+ class_weights=class_weights,
740
+ )
741
+ result = trainer.train()
742
+ assert np.isfinite(result.training_loss)
@@ -1,17 +0,0 @@
1
- """bertuner: hyperparameter optimization and fine-tuning for BERT-style text classifiers."""
2
-
3
- __version__ = "0.1.0"
4
-
5
- from bertuner.Predictor import BERTunePredictor
6
-
7
- __all__ = ["BERTuneClassifier", "BERTunePredictor", "__version__"]
8
-
9
-
10
- def __getattr__(name):
11
- # Lazy import: BERTuneClassifier pulls training-only deps (optuna, mlflow,
12
- # datasets), which are optional extras — inference installs must not need them.
13
- if name == "BERTuneClassifier":
14
- from bertuner.BERTuner import BERTuneClassifier
15
-
16
- return BERTuneClassifier
17
- raise AttributeError(f"module 'bertuner' has no attribute '{name}'")
File without changes
File without changes
File without changes
File without changes
File without changes