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.
- {bertuner-0.1.0 → bertuner-0.1.2}/PKG-INFO +22 -9
- {bertuner-0.1.0 → bertuner-0.1.2}/README.md +14 -3
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/BERTuner.py +264 -34
- bertuner-0.1.2/bertuner/__init__.py +8 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/constants.py +0 -2
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner.egg-info/PKG-INFO +22 -9
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner.egg-info/requires.txt +7 -5
- {bertuner-0.1.0 → bertuner-0.1.2}/pyproject.toml +10 -6
- {bertuner-0.1.0 → bertuner-0.1.2}/tests/test_bertuner.py +255 -0
- bertuner-0.1.0/bertuner/__init__.py +0 -17
- {bertuner-0.1.0 → bertuner-0.1.2}/LICENSE +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/CustomTrainer.py +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/Predictor.py +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/TensorBoardCallback.py +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner/utils.py +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner.egg-info/SOURCES.txt +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner.egg-info/dependency_links.txt +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/bertuner.egg-info/top_level.txt +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/setup.cfg +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/tests/test_predictor.py +0 -0
- {bertuner-0.1.0 → bertuner-0.1.2}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: bertuner
|
|
3
|
-
Version: 0.1.
|
|
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).
|
|
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
|
|
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).
|
|
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
|
|
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
|
-
|
|
214
|
-
|
|
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
|
-
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
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
|
-
|
|
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,
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
646
|
-
|
|
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
|
-
"
|
|
731
|
-
|
|
732
|
-
|
|
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
|
-
"
|
|
747
|
-
"
|
|
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
|
-
"
|
|
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
|
-
"
|
|
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.
|
|
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).
|
|
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
|
|
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.
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|