bertuner 0.2.3__tar.gz → 0.2.4__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.
Files changed (24) hide show
  1. {bertuner-0.2.3 → bertuner-0.2.4}/PKG-INFO +5 -1
  2. {bertuner-0.2.3 → bertuner-0.2.4}/README.md +4 -0
  3. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/BERTuner.py +72 -26
  4. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/__init__.py +1 -1
  5. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/PKG-INFO +5 -1
  6. {bertuner-0.2.3 → bertuner-0.2.4}/pyproject.toml +1 -1
  7. {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_bertuner.py +81 -0
  8. {bertuner-0.2.3 → bertuner-0.2.4}/LICENSE +0 -0
  9. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/CustomTrainer.py +0 -0
  10. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/Predictor.py +0 -0
  11. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/TensorBoardCallback.py +0 -0
  12. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/compat.py +0 -0
  13. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/constants.py +0 -0
  14. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/exceptions.py +0 -0
  15. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/utils.py +0 -0
  16. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/SOURCES.txt +0 -0
  17. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/dependency_links.txt +0 -0
  18. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/requires.txt +0 -0
  19. {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/top_level.txt +0 -0
  20. {bertuner-0.2.3 → bertuner-0.2.4}/setup.cfg +0 -0
  21. {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_numerical_stability.py +0 -0
  22. {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_predictor.py +0 -0
  23. {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_transformers_compat.py +0 -0
  24. {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.2.3
3
+ Version: 0.2.4
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
@@ -121,6 +121,10 @@ print(metrics)
121
121
 
122
122
  Multi-label classification: pass several target columns — `target_cols=["l1", "l2", "l3"]`. The loss switches to BCE-with-logits and one decision threshold is optimised per label.
123
123
 
124
+ Pass `threshold_metric="balanced_accuracy"` to `BERTuneClassifier` to choose thresholds by balanced accuracy. The default, `threshold_metric="f1"`, optimizes F-beta with `threshold_beta=1.0` (F1). Set `threshold_beta=2.0` for F2, which gives recall more weight, or `threshold_beta=0.5` for F0.5, which gives precision more weight. Beta must be positive and finite; it is ignored when using balanced accuracy.
125
+
126
+ Threshold tuning runs during `train_final_model()` using validation predictions only and is independent of `optimize_metric`, which selects the model during Optuna optimization. It checks every distinct predicted probability, plus 0.5 and a boundary above the maximum for all-negative predictions. Ties prefer 0.5, then the lowest optimal threshold. Multi-label thresholds are chosen independently for each label; multiclass predictions use argmax. The saved `bertuner_config.json` includes `threshold_metric` and `threshold_beta` alongside `optimal_threshold`, and the predictor automatically uses the saved threshold. Reported F1 metrics remain F1 even when the threshold is tuned for another beta.
127
+
124
128
  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.
125
129
 
126
130
  ### Using existing splits
@@ -80,6 +80,10 @@ print(metrics)
80
80
 
81
81
  Multi-label classification: pass several target columns — `target_cols=["l1", "l2", "l3"]`. The loss switches to BCE-with-logits and one decision threshold is optimised per label.
82
82
 
83
+ Pass `threshold_metric="balanced_accuracy"` to `BERTuneClassifier` to choose thresholds by balanced accuracy. The default, `threshold_metric="f1"`, optimizes F-beta with `threshold_beta=1.0` (F1). Set `threshold_beta=2.0` for F2, which gives recall more weight, or `threshold_beta=0.5` for F0.5, which gives precision more weight. Beta must be positive and finite; it is ignored when using balanced accuracy.
84
+
85
+ Threshold tuning runs during `train_final_model()` using validation predictions only and is independent of `optimize_metric`, which selects the model during Optuna optimization. It checks every distinct predicted probability, plus 0.5 and a boundary above the maximum for all-negative predictions. Ties prefer 0.5, then the lowest optimal threshold. Multi-label thresholds are chosen independently for each label; multiclass predictions use argmax. The saved `bertuner_config.json` includes `threshold_metric` and `threshold_beta` alongside `optimal_threshold`, and the predictor automatically uses the saved threshold. Reported F1 metrics remain F1 even when the threshold is tuned for another beta.
86
+
83
87
  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.
84
88
 
85
89
  ### Using existing splits
@@ -3,6 +3,7 @@ import random
3
3
  import json
4
4
  import shutil
5
5
  import warnings
6
+ from numbers import Real
6
7
  import numpy as np
7
8
  import pandas as pd
8
9
  import torch
@@ -89,11 +90,21 @@ class BERTuneClassifier:
89
90
  max_grad_norm: float | None = 1.0,
90
91
  class_weight_warning_threshold: float | None = 100.0,
91
92
  data_splits: dict[str, pd.DataFrame | str | os.PathLike] = None,
93
+ threshold_metric: str = "f1",
94
+ threshold_beta: float = 1.0,
92
95
  ):
93
96
  if sum(source is not None for source in (data_path, dataframe, data_splits)) != 1:
94
97
  raise ValueError("Provide exactly one of data_path (CSV), dataframe, or data_splits.")
95
98
  if precision not in {"auto", "fp32", "bf16", "fp16"}:
96
99
  raise ValueError("precision must be one of: 'auto', 'fp32', 'bf16', 'fp16'.")
100
+ if threshold_metric not in ("f1", "balanced_accuracy"):
101
+ raise ValueError("threshold_metric must be 'f1' or 'balanced_accuracy'.")
102
+ if (
103
+ not isinstance(threshold_beta, Real)
104
+ or not np.isfinite(threshold_beta)
105
+ or threshold_beta <= 0
106
+ ):
107
+ raise ValueError("threshold_beta must be a positive finite number.")
97
108
  if max_grad_norm is not None and (
98
109
  not np.isfinite(max_grad_norm) or max_grad_norm <= 0
99
110
  ):
@@ -156,6 +167,8 @@ class BERTuneClassifier:
156
167
  self.best_precision_fallback = False
157
168
  # Single-label: scalar float. Multi-label: array of per-label floats.
158
169
  self.best_threshold = 0.5
170
+ self.threshold_metric = threshold_metric
171
+ self.threshold_beta = float(threshold_beta)
159
172
  self.max_length = max_length
160
173
  # None → auto: enabled when the effective sequence length is long enough
161
174
  # that activation memory dominates (see _use_gradient_checkpointing).
@@ -1219,10 +1232,14 @@ class BERTuneClassifier:
1219
1232
  """
1220
1233
  Finds the best classification threshold(s) on the validation set.
1221
1234
 
1222
- Binary → one scalar threshold (maximises F1).
1235
+ Binary → one scalar threshold.
1223
1236
  Multiclass → None (predictions are argmax; thresholds don't apply).
1224
- Multi-label → one threshold per label (maximises macro-F1);
1237
+ Multi-label → one independently optimised threshold per label;
1225
1238
  returns np.ndarray of shape (num_labels,).
1239
+
1240
+ Maximises F-beta (threshold_metric='f1') or balanced accuracy.
1241
+ Evaluates every distinct prediction boundary, including all-negative
1242
+ predictions. Ties prefer 0.5, then the lowest optimal threshold.
1226
1243
  """
1227
1244
  if not self.is_multilabel and not self.is_binary:
1228
1245
  return None
@@ -1231,31 +1248,58 @@ class BERTuneClassifier:
1231
1248
  labels = val_res.label_ids
1232
1249
 
1233
1250
  if self.is_multilabel:
1234
- # Optimise each label independently
1235
- best_thresholds = np.full(self.num_labels, 0.5)
1236
- for i in range(self.num_labels):
1237
- best_f1, best_t = 0.0, 0.5
1238
- for thresh in np.linspace(0.1, 0.9, 81):
1239
- f1 = f1_score(
1240
- labels[:, i],
1241
- (probs[:, i] >= thresh).astype(int),
1242
- zero_division=0,
1243
- )
1244
- if f1 > best_f1:
1245
- best_f1, best_t = f1, thresh
1246
- best_thresholds[i] = best_t
1247
- return best_thresholds
1251
+ return np.array([
1252
+ self._best_binary_threshold(labels[:, i], probs[:, i])
1253
+ for i in range(self.num_labels)
1254
+ ])
1255
+ return self._best_binary_threshold(labels, probs)
1256
+
1257
+ def _best_binary_threshold(self, labels, probs):
1258
+ """Score distinct boundaries in O(n log n), with tied scores kept together."""
1259
+ order = np.argsort(probs)
1260
+ sorted_probs = probs[order]
1261
+ positives = np.concatenate(([0], np.cumsum(labels[order] == 1)))
1262
+ # Use the probability dtype so the all-negative boundary remains above
1263
+ # the maximum when inference compares against float32 probabilities.
1264
+ above_max = float(np.nextafter(probs.max(), np.array(np.inf, dtype=probs.dtype)))
1265
+ thresholds = np.unique(np.append(sorted_probs, [0.5, above_max]))
1266
+ boundaries = np.searchsorted(sorted_probs, thresholds, side="left")
1267
+ tp = positives[-1] - positives[boundaries]
1268
+ fn = positives[boundaries]
1269
+ fp = len(labels) - boundaries - tp
1270
+ tn = boundaries - fn
1271
+
1272
+ if self.threshold_metric == "balanced_accuracy":
1273
+ scores = np.zeros(len(thresholds))
1274
+ present_classes = 0
1275
+ if positives[-1] > 0:
1276
+ scores += tp / positives[-1]
1277
+ present_classes += 1
1278
+ negatives = len(labels) - positives[-1]
1279
+ if negatives > 0:
1280
+ scores += tn / negatives
1281
+ present_classes += 1
1282
+ scores /= present_classes
1248
1283
  else:
1249
- best_f1, best_t = 0.0, 0.5
1250
- for thresh in np.linspace(0.1, 0.9, 81):
1251
- f1 = f1_score(
1252
- labels,
1253
- (probs >= thresh).astype(int),
1254
- zero_division=0,
1255
- )
1256
- if f1 > best_f1:
1257
- best_f1, best_t = f1, thresh
1258
- return best_t
1284
+ # Normalised beta weights avoid overflow for large finite beta.
1285
+ beta = self.threshold_beta
1286
+ if beta <= 1:
1287
+ recall_weight = beta * beta / (1 + beta * beta)
1288
+ precision_weight = 1 / (1 + beta * beta)
1289
+ else:
1290
+ inverse_square = (1 / beta) ** 2
1291
+ recall_weight = 1 / (1 + inverse_square)
1292
+ precision_weight = inverse_square / (1 + inverse_square)
1293
+ denominator = tp + recall_weight * fn + precision_weight * fp
1294
+ scores = np.divide(
1295
+ tp, denominator, out=np.zeros(len(thresholds)), where=denominator > 0
1296
+ )
1297
+
1298
+ best_score = scores.max()
1299
+ default_index = np.searchsorted(thresholds, 0.5)
1300
+ if scores[default_index] == best_score:
1301
+ return 0.5
1302
+ return float(thresholds[np.argmax(scores)])
1259
1303
 
1260
1304
  # ------------------------------------------------------------------
1261
1305
  # Metrics DataFrame
@@ -1404,6 +1448,8 @@ class BERTuneClassifier:
1404
1448
  "model": self.best_params["model"],
1405
1449
  "model_path": model_path,
1406
1450
  "optimal_threshold": threshold,
1451
+ "threshold_metric": self.threshold_metric,
1452
+ "threshold_beta": self.threshold_beta,
1407
1453
  "is_multilabel": self.is_multilabel,
1408
1454
  "target_cols": self.target_cols,
1409
1455
  "max_length": max_length if max_length is not None else self.max_length,
@@ -1,6 +1,6 @@
1
1
  """bertuner: hyperparameter optimization and fine-tuning for BERT-style text classifiers."""
2
2
 
3
- __version__ = "0.2.3"
3
+ __version__ = "0.2.4"
4
4
 
5
5
  from bertuner.BERTuner import BERTuneClassifier
6
6
  from bertuner.Predictor import BERTunePredictor
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.2.3
3
+ Version: 0.2.4
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
@@ -121,6 +121,10 @@ print(metrics)
121
121
 
122
122
  Multi-label classification: pass several target columns — `target_cols=["l1", "l2", "l3"]`. The loss switches to BCE-with-logits and one decision threshold is optimised per label.
123
123
 
124
+ Pass `threshold_metric="balanced_accuracy"` to `BERTuneClassifier` to choose thresholds by balanced accuracy. The default, `threshold_metric="f1"`, optimizes F-beta with `threshold_beta=1.0` (F1). Set `threshold_beta=2.0` for F2, which gives recall more weight, or `threshold_beta=0.5` for F0.5, which gives precision more weight. Beta must be positive and finite; it is ignored when using balanced accuracy.
125
+
126
+ Threshold tuning runs during `train_final_model()` using validation predictions only and is independent of `optimize_metric`, which selects the model during Optuna optimization. It checks every distinct predicted probability, plus 0.5 and a boundary above the maximum for all-negative predictions. Ties prefer 0.5, then the lowest optimal threshold. Multi-label thresholds are chosen independently for each label; multiclass predictions use argmax. The saved `bertuner_config.json` includes `threshold_metric` and `threshold_beta` alongside `optimal_threshold`, and the predictor automatically uses the saved threshold. Reported F1 metrics remain F1 even when the threshold is tuned for another beta.
127
+
124
128
  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.
125
129
 
126
130
  ### Using existing splits
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "bertuner"
7
- version = "0.2.3"
7
+ version = "0.2.4"
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" }
@@ -9,6 +9,7 @@ import pandas as pd
9
9
  import pytest
10
10
  import optuna
11
11
  import mlflow
12
+ from sklearn.metrics import balanced_accuracy_score, fbeta_score
12
13
  from mlflow.tracking import MlflowClient
13
14
 
14
15
  from bertuner.BERTuner import BERTuneClassifier
@@ -244,6 +245,75 @@ class TestGetProbs:
244
245
 
245
246
 
246
247
  class TestOptimizeThreshold:
248
+ @pytest.mark.parametrize("metric", ["auc_roc", "fbeta", None])
249
+ def test_invalid_metric(self, tmp_path, metric):
250
+ with pytest.raises(ValueError, match="threshold_metric"):
251
+ make_classifier(tmp_path, threshold_metric=metric)
252
+
253
+ @pytest.mark.parametrize("beta", [0, -1, np.nan, np.inf, "2", None])
254
+ def test_invalid_beta(self, tmp_path, beta):
255
+ with pytest.raises(ValueError, match="threshold_beta"):
256
+ make_classifier(tmp_path, threshold_beta=beta)
257
+
258
+ @pytest.mark.parametrize("metric", ["f1", "balanced_accuracy"])
259
+ def test_clustered_probabilities(self, tmp_path, monkeypatch, metric):
260
+ clf = make_classifier(tmp_path, threshold_metric=metric)
261
+ probs = np.array([0.531, 0.534, 0.536, 0.539])
262
+ monkeypatch.setattr(clf, "_get_probs", lambda _: probs)
263
+ threshold = clf._optimize_threshold(SimpleNamespace(
264
+ predictions=None, label_ids=np.array([0, 0, 1, 1])
265
+ ))
266
+ assert threshold == 0.536
267
+
268
+ @pytest.mark.parametrize("metric, beta, expected", [
269
+ ("f1", 0.5, 0.9), ("f1", 1, 0.5), ("f1", 2, 0.5),
270
+ ("balanced_accuracy", 1, 0.9),
271
+ ])
272
+ def test_metric_changes_optimal_threshold(self, tmp_path, monkeypatch, metric, beta, expected):
273
+ clf = make_classifier(tmp_path, threshold_metric=metric, threshold_beta=beta)
274
+ probs = np.array([0.9, 0.8, 0.7, 0.6, 0.55, 0.54])
275
+ monkeypatch.setattr(clf, "_get_probs", lambda _: probs)
276
+ threshold = clf._optimize_threshold(SimpleNamespace(
277
+ predictions=None, label_ids=np.array([1, 0, 0, 1, 0, 1])
278
+ ))
279
+ assert threshold == expected
280
+
281
+ @pytest.mark.parametrize("metric, beta", [
282
+ ("f1", 1), ("f1", 2), ("f1", 0.5), ("balanced_accuracy", 1)
283
+ ])
284
+ def test_matches_exhaustive_search(self, tmp_path, metric, beta):
285
+ clf = make_classifier(tmp_path, threshold_metric=metric, threshold_beta=beta)
286
+ rng = np.random.default_rng(17)
287
+ for _ in range(10):
288
+ probs = rng.choice([0.02, 0.531, 0.534, 0.539, 0.98], 30)
289
+ labels = rng.integers(0, 2, len(probs))
290
+ candidates = np.unique(np.append(probs, [0.5, 1.0]))
291
+ def score(t):
292
+ preds = probs >= t
293
+ if metric == "balanced_accuracy":
294
+ return balanced_accuracy_score(labels, preds)
295
+ return fbeta_score(labels, preds, beta=beta, zero_division=0)
296
+ threshold = clf._best_binary_threshold(labels, probs)
297
+ assert score(threshold) == pytest.approx(max(map(score, candidates)))
298
+
299
+ @pytest.mark.parametrize("metric", ["f1", "balanced_accuracy"])
300
+ def test_multilabel_single_class_and_saturated_scores(self, tmp_path, monkeypatch, metric):
301
+ cols = ["l1", "l2"]
302
+ clf = make_classifier(
303
+ tmp_path, target_cols=cols, dataframe=make_df(multilabel_cols=cols),
304
+ threshold_metric=metric,
305
+ )
306
+ probs = np.array([[1, 0], [1, 0]], dtype=np.float32)
307
+ labels = np.array([[0, 1], [0, 1]])
308
+ monkeypatch.setattr(clf, "_get_probs", lambda _: probs)
309
+ thresholds = clf._optimize_threshold(SimpleNamespace(predictions=None, label_ids=labels))
310
+ assert thresholds[1] == 0
311
+ if metric == "balanced_accuracy":
312
+ assert thresholds[0] > 1
313
+ assert not np.any(probs[:, 0] >= thresholds[0])
314
+ else:
315
+ assert thresholds[0] == 0.5 # All F-beta scores are zero.
316
+
247
317
  def test_single_label_finds_separating_threshold(self, tmp_path):
248
318
  clf = make_classifier(tmp_path)
249
319
  # Positives cluster at high prob, negatives at low → best threshold between
@@ -529,11 +599,22 @@ class TestSaveModel:
529
599
  meta = config["model_metadata"]
530
600
  assert meta["model_path"] == "bert-base-uncased"
531
601
  assert meta["optimal_threshold"] == 0.42
602
+ assert meta["threshold_metric"] == "f1"
603
+ assert meta["threshold_beta"] == 1.0
532
604
  assert meta["is_multilabel"] is False
533
605
  assert meta["target_cols"] == ["target"]
534
606
  assert meta["max_length"] == 512
535
607
  assert config["parameters"] == clf.best_params
536
608
 
609
+ @pytest.mark.parametrize("metric, beta", [("balanced_accuracy", 1), ("f1", 2)])
610
+ def test_saves_threshold_settings(self, tmp_path, metric, beta):
611
+ clf = make_classifier(tmp_path, threshold_metric=metric, threshold_beta=beta)
612
+ clf.best_params = {"model": "bert-base"}
613
+ clf._save_model(str(tmp_path), MagicMock(), MagicMock(), "bert-base-uncased")
614
+ config = json.loads((tmp_path / "model" / "bertuner_config.json").read_text())
615
+ assert config["model_metadata"]["threshold_metric"] == metric
616
+ assert config["model_metadata"]["threshold_beta"] == beta
617
+
537
618
  def test_multilabel_threshold_serialised_as_list(self, tmp_path):
538
619
  cols = ["l1", "l2"]
539
620
  clf = make_classifier(
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes