BERTuneClassifier
A library for hyperparameter optimization and fine-tuning of BERT-based classification models. It integrates Optuna for efficient search and MLflow for experiment tracking.
Supports both classic 512-token encoders (BERT, RoBERTa, DistilBERT, ELECTRA) and long-context models such as ModernBERT (8192 tokens).
Installation
pip install bertuner # training + inference, batteries included
From source (development):
git clone https://github.com/elemets/bertuner && cd bertuner
pip install -r requirements.txt
MLflow tracking works in two modes:
# Option A: run a tracking server (default, expects port 9090)
mlflow server --port 9090
# Option B: no server — log to a local directory instead
classifier = BERTuneClassifier(..., mlflow_tracking_uri="./mlruns")
Training
from bertuner.BERTuner import BERTuneClassifier
# 1. Initialize
classifier = BERTuneClassifier(
data_path="../data/dataset.csv", # or dataframe=my_df
models_dir="../models/",
text_feature="text_col", # column containing the text
target_cols=["label_col"], # one column = single-label
max_length=512,
precision="auto", # BF16/FP16 on supported CUDA, else FP32
retry_nonfinite_in_fp32=True, # retry verified NaN/Inf failures once
max_grad_norm=1.0, # explicit gradient clipping threshold
)
# 2. Configure (optional: uses defaults if called without arguments)
classifier.initialize_model_choices()
classifier.initialize_search_space()
# 3. Optimize — runs Optuna trials and logs to MLflow
best_value = classifier.optimize(
n_trials=20,
optimize_metric="avg_precision",
study_name="bert_experiment_v1",
)
# 4. Train final model — retrains on best params, optimises the decision
# threshold on the validation set, evaluates on the test set, and saves
# model + tokenizer + bertuner_config.json under models_dir/final_model/model
metrics, model, test_ds = classifier.train_final_model()
print(metrics)
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.
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.
Numerical-stability recovery
Every attempt checks logits, loss, gradients before the optimizer step, and evaluation predictions for NaN/Inf. When a mixed-precision attempt becomes non-finite, BERTuner deletes its checkpoints, reloads the pretrained model, resets the seed, and retries the same Optuna trial and hyperparameters once in FP32. A second numerical failure prunes an optimization trial; final-model training raises NonFiniteTrainingError. Invalid labels, class weights, OOM errors, and other exceptions never trigger the retry.
precision accepts "auto", "fp32", "bf16", or "fp16". Explicit unsupported precision raises before training. Set retry_nonfinite_in_fp32=False to disable recovery, max_grad_norm=None to disable clipping, or class_weight_warning_threshold=None to disable warnings for large finite class weights. Saved bertuner_config.json records requested/effective precision and fallback status.
Metrics logged to MLflow
Final runs log one canonical metric set for both Validation_* and Test_*:
- Binary: accuracy, balanced accuracy, precision, recall, specificity, F1, Matthews correlation coefficient (MCC), average precision, AUROC, Brier score, and log loss.
- Multiclass: accuracy, balanced accuracy, macro and weighted precision/recall/F1, MCC, macro and weighted average precision/AUROC, and log loss.
- 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.
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.
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.
Customizing the hyperparameter search
Two things are configurable: which models are searched and which hyperparameters with what ranges.
initialize_model_choices maps short names to HuggingFace model paths:
classifier.initialize_model_choices({
"bert-base": "bert-base-uncased",
"modernbert-base": "answerdotai/ModernBERT-base",
"my-domain-model": "allenai/scibert_scivocab_uncased",
})
initialize_search_space takes a dict where the value type decides the Optuna suggestion:
- list → categorical choice, e.g.
"batch_size": [8, 16, 32] - dict with int
low/high→ integer range, e.g.{"low": 3, "high": 8}(optional"step") - dict with float
low/high→ float range, e.g.{"low": 1e-6, "high": 5e-5, "log": True}("log"samples on a log scale — use it for learning rates)
classifier.initialize_search_space({
"model": ["bert-base", "my-domain-model"], # keys from model_choices
"learning_rate": {"low": 1e-6, "high": 5e-5, "log": True},
"batch_size": [8, 16, 32],
"gradient_accumulation_steps": [1, 2, 4], # optional, defaults to 1
"loss_type": ["weighted", "focal", "label_smoothing"],
"weight_decay": {"low": 0.0, "high": 0.2},
"warmup_ratio": {"low": 0.0, "high": 0.2},
"scheduler": ["linear", "cosine"],
"dropout": {"low": 0.0, "high": 0.3},
"early_stopping_patience": {"low": 3, "high": 8},
})
Required keys: model, learning_rate, batch_size, weight_decay, warmup_ratio, scheduler, dropout, early_stopping_patience. Optional: loss_type (single-label only; defaults to weighted) and gradient_accumulation_steps.
Ready-made spaces live in bertuner.constants: DEFAULT_SEARCH_SPACE_SINGLELABEL, DEFAULT_SEARCH_SPACE_MULTILABEL, and DEFAULT_SEARCH_SPACE_LONGCONTEXT. Tweak one instead of starting from scratch:
from bertuner.constants import DEFAULT_SEARCH_SPACE_SINGLELABEL
classifier.initialize_search_space({
**DEFAULT_SEARCH_SPACE_SINGLELABEL,
"model": ["bert-base"], # pin a single model
"learning_rate": {"low": 1e-5, "high": 3e-5, "log": True},
})
Long documents (ModernBERT, 8192 tokens)
from bertuner.BERTuner import BERTuneClassifier
from bertuner.constants import DEFAULT_SEARCH_SPACE_LONGCONTEXT
classifier = BERTuneClassifier(
data_path="../data/long_docs.csv",
models_dir="../models/",
text_feature="text_col",
target_cols=["label_col"],
max_length=8192,
)
classifier.initialize_model_choices()
classifier.initialize_search_space(DEFAULT_SEARCH_SPACE_LONGCONTEXT)
classifier.optimize(n_trials=10, study_name="long_context_v1")
DEFAULT_SEARCH_SPACE_LONGCONTEXT searches over ModernBERT base/large with small per-device batches and gradient_accumulation_steps, keeping the effective batch size in the usual range without exhausting GPU memory. Mixing 512-token models into the same search space is safe — max_length is clamped per model.
Loading a trained model and predicting
train_final_model() saves everything the predictor needs (weights, tokenizer, optimised thresholds, max_length) under models_dir/final_model/model:
from bertuner.Predictor import BERTunePredictor
predictor = BERTunePredictor("../models/final_model/model")
# Hard class predictions, using the threshold(s) optimised during training
preds = predictor.predict(["some clinical note", "another document"])
# single-label → array of 0/1 (binary) or class ids (multiclass)
# multi-label → array of shape (N, num_labels) with 0/1 per label
# Probabilities
probs = predictor.predict_proba(["some clinical note"])
# single-label → softmax over classes, shape (N, num_classes)
# multi-label → sigmoid per label, shape (N, num_labels)
# Predictions as a DataFrame with one column per target
df = predictor.predict_df(["some clinical note", "another document"])
Options: BERTunePredictor(model_dir, device="cuda", batch_size=64) — device defaults to CUDA when available, batch size to 32. Texts longer than the trained max_length are truncated.
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file bertuner-0.2.1.tar.gz.
File metadata
- Download URL: bertuner-0.2.1.tar.gz
- Upload date:
- Size: 38.0 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
c70c05fd7ac78ca57098e85a62d5911554e62a1247b522e9a738910ebeb3a19a
|
|
| MD5 |
bc42fbfdd40f5b65ebd61c3d068345cd
|
|
| BLAKE2b-256 |
5b496e8b2a62521c1e8f5f99e740bc170f8b5cadd2cca2b0d9ad0f106fbf8bc0
|
Provenance
The following attestation bundles were made for bertuner-0.2.1.tar.gz:
Publisher:
workflow.yml on elemets/bertuner
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
bertuner-0.2.1.tar.gz -
Subject digest:
c70c05fd7ac78ca57098e85a62d5911554e62a1247b522e9a738910ebeb3a19a - Sigstore transparency entry: 2227814789
- Sigstore integration time:
-
Permalink:
elemets/bertuner@c9e956d9c638dc227091f542402443c3328cd137 -
Branch / Tag:
refs/tags/v0.2.1 - Owner: https://github.com/elemets
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
workflow.yml@c9e956d9c638dc227091f542402443c3328cd137 -
Trigger Event:
release
-
Statement type:
File details
Details for the file bertuner-0.2.1-py3-none-any.whl.
File metadata
- Download URL: bertuner-0.2.1-py3-none-any.whl
- Upload date:
- Size: 27.1 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
0640c93a03862c12722614bcbae0b5b6b8bfa38990ebebf9ddb666cfd015f741
|
|
| MD5 |
f494ad4843e8b9ef3cc17a67ecdbd5e6
|
|
| BLAKE2b-256 |
aa2aecb9546087710219460320e8eaceba91e76821c4e3a53dcf6b3048f59db2
|
Provenance
The following attestation bundles were made for bertuner-0.2.1-py3-none-any.whl:
Publisher:
workflow.yml on elemets/bertuner
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
bertuner-0.2.1-py3-none-any.whl -
Subject digest:
0640c93a03862c12722614bcbae0b5b6b8bfa38990ebebf9ddb666cfd015f741 - Sigstore transparency entry: 2227814932
- Sigstore integration time:
-
Permalink:
elemets/bertuner@c9e956d9c638dc227091f542402443c3328cd137 -
Branch / Tag:
refs/tags/v0.2.1 - Owner: https://github.com/elemets
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
workflow.yml@c9e956d9c638dc227091f542402443c3328cd137 -
Trigger Event:
release
-
Statement type: