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.
- {bertuner-0.2.3 → bertuner-0.2.4}/PKG-INFO +5 -1
- {bertuner-0.2.3 → bertuner-0.2.4}/README.md +4 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/BERTuner.py +72 -26
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/__init__.py +1 -1
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/PKG-INFO +5 -1
- {bertuner-0.2.3 → bertuner-0.2.4}/pyproject.toml +1 -1
- {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_bertuner.py +81 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/LICENSE +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/CustomTrainer.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/Predictor.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/TensorBoardCallback.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/compat.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/constants.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/exceptions.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner/utils.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/SOURCES.txt +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/dependency_links.txt +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/requires.txt +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/bertuner.egg-info/top_level.txt +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/setup.cfg +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_numerical_stability.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_predictor.py +0 -0
- {bertuner-0.2.3 → bertuner-0.2.4}/tests/test_transformers_compat.py +0 -0
- {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
|
+
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
|
|
1235
|
+
Binary → one scalar threshold.
|
|
1223
1236
|
Multiclass → None (predictions are argmax; thresholds don't apply).
|
|
1224
|
-
Multi-label → one threshold per label
|
|
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
|
-
|
|
1235
|
-
|
|
1236
|
-
|
|
1237
|
-
|
|
1238
|
-
|
|
1239
|
-
|
|
1240
|
-
|
|
1241
|
-
|
|
1242
|
-
|
|
1243
|
-
|
|
1244
|
-
|
|
1245
|
-
|
|
1246
|
-
|
|
1247
|
-
|
|
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
|
-
|
|
1250
|
-
|
|
1251
|
-
|
|
1252
|
-
|
|
1253
|
-
|
|
1254
|
-
|
|
1255
|
-
)
|
|
1256
|
-
|
|
1257
|
-
|
|
1258
|
-
|
|
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
|
Metadata-Version: 2.4
|
|
2
2
|
Name: bertuner
|
|
3
|
-
Version: 0.2.
|
|
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.
|
|
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
|
|
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
|