bertuner 0.2.1__tar.gz → 0.2.3__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.1 → bertuner-0.2.3}/PKG-INFO +36 -1
  2. {bertuner-0.2.1 → bertuner-0.2.3}/README.md +35 -0
  3. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/BERTuner.py +81 -20
  4. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/CustomTrainer.py +3 -0
  5. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/__init__.py +1 -1
  6. bertuner-0.2.3/bertuner/compat.py +15 -0
  7. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner.egg-info/PKG-INFO +36 -1
  8. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner.egg-info/SOURCES.txt +2 -0
  9. {bertuner-0.2.1 → bertuner-0.2.3}/pyproject.toml +1 -1
  10. {bertuner-0.2.1 → bertuner-0.2.3}/tests/test_bertuner.py +90 -2
  11. bertuner-0.2.3/tests/test_transformers_compat.py +88 -0
  12. {bertuner-0.2.1 → bertuner-0.2.3}/LICENSE +0 -0
  13. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/Predictor.py +0 -0
  14. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/TensorBoardCallback.py +0 -0
  15. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/constants.py +0 -0
  16. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/exceptions.py +0 -0
  17. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner/utils.py +0 -0
  18. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner.egg-info/dependency_links.txt +0 -0
  19. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner.egg-info/requires.txt +0 -0
  20. {bertuner-0.2.1 → bertuner-0.2.3}/bertuner.egg-info/top_level.txt +0 -0
  21. {bertuner-0.2.1 → bertuner-0.2.3}/setup.cfg +0 -0
  22. {bertuner-0.2.1 → bertuner-0.2.3}/tests/test_numerical_stability.py +0 -0
  23. {bertuner-0.2.1 → bertuner-0.2.3}/tests/test_predictor.py +0 -0
  24. {bertuner-0.2.1 → bertuner-0.2.3}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.2.1
3
+ Version: 0.2.3
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
@@ -58,6 +58,20 @@ git clone https://github.com/elemets/bertuner && cd bertuner
58
58
  pip install -r requirements.txt
59
59
  ```
60
60
 
61
+ BERTuner supports Transformers 4.48+ and 5.x. Compatibility tests cover the
62
+ 4.48 baseline and the latest available 4.x and 5.x releases, including final
63
+ training, TensorBoard logging, warmup scheduling, and gradient accumulation.
64
+ Future releases are checked by a weekly CI run rather than assumed compatible.
65
+
66
+ When initializing a classifier from a base pretrained model, a Transformers 5
67
+ load report may list `lm_head.*` weights as `UNEXPECTED` and classifier weights
68
+ as `MISSING`. This is expected when replacing the language-model head with a
69
+ classification head; the new head is learned during fine-tuning. Unexpected
70
+ encoder weights or shape mismatches should still be investigated.
71
+
72
+ Run the test suite with `python -m pytest tests -q`. Training regression tests
73
+ use a tiny local checkpoint; cached-model predictor tests skip if unavailable.
74
+
61
75
  MLflow tracking works in two modes:
62
76
 
63
77
  ```bash
@@ -109,6 +123,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
109
123
 
110
124
  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.
111
125
 
126
+ ### Using existing splits
127
+
128
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
129
+
130
+ ```python
131
+ classifier = BERTuneClassifier(
132
+ models_dir="../models/",
133
+ text_feature="text_col",
134
+ target_cols=["label_col"],
135
+ data_splits={
136
+ "train": train_df,
137
+ "validation": validation_df, # "val" also accepted
138
+ "test": test_df,
139
+ },
140
+ )
141
+ ```
142
+
143
+ Each value can also be a CSV path, such as `"data/train.csv"` (including `pathlib.Path` objects). All three splits must be nonempty and contain the text and target columns. Use the same numeric label encoding across all splits. Both optimization and final training preserve these partitions and their row order; no random splitting or stratification is performed. Class weights are calculated from training data, and threshold selection uses validation data.
144
+
145
+ If you supply `group_key`, each split must contain that column with nonmissing IDs. Overlapping groups across splits raise an error, including in multi-label mode. Group comparisons ignore surrounding whitespace and case. Without `group_key`, checking for overlap between external partitions is the caller's responsibility; DataFrame indices are not treated as row IDs.
146
+
112
147
  ### Numerical-stability recovery
113
148
 
114
149
  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.
@@ -17,6 +17,20 @@ git clone https://github.com/elemets/bertuner && cd bertuner
17
17
  pip install -r requirements.txt
18
18
  ```
19
19
 
20
+ BERTuner supports Transformers 4.48+ and 5.x. Compatibility tests cover the
21
+ 4.48 baseline and the latest available 4.x and 5.x releases, including final
22
+ training, TensorBoard logging, warmup scheduling, and gradient accumulation.
23
+ Future releases are checked by a weekly CI run rather than assumed compatible.
24
+
25
+ When initializing a classifier from a base pretrained model, a Transformers 5
26
+ load report may list `lm_head.*` weights as `UNEXPECTED` and classifier weights
27
+ as `MISSING`. This is expected when replacing the language-model head with a
28
+ classification head; the new head is learned during fine-tuning. Unexpected
29
+ encoder weights or shape mismatches should still be investigated.
30
+
31
+ Run the test suite with `python -m pytest tests -q`. Training regression tests
32
+ use a tiny local checkpoint; cached-model predictor tests skip if unavailable.
33
+
20
34
  MLflow tracking works in two modes:
21
35
 
22
36
  ```bash
@@ -68,6 +82,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
68
82
 
69
83
  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.
70
84
 
85
+ ### Using existing splits
86
+
87
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
88
+
89
+ ```python
90
+ classifier = BERTuneClassifier(
91
+ models_dir="../models/",
92
+ text_feature="text_col",
93
+ target_cols=["label_col"],
94
+ data_splits={
95
+ "train": train_df,
96
+ "validation": validation_df, # "val" also accepted
97
+ "test": test_df,
98
+ },
99
+ )
100
+ ```
101
+
102
+ Each value can also be a CSV path, such as `"data/train.csv"` (including `pathlib.Path` objects). All three splits must be nonempty and contain the text and target columns. Use the same numeric label encoding across all splits. Both optimization and final training preserve these partitions and their row order; no random splitting or stratification is performed. Class weights are calculated from training data, and threshold selection uses validation data.
103
+
104
+ If you supply `group_key`, each split must contain that column with nonmissing IDs. Overlapping groups across splits raise an error, including in multi-label mode. Group comparisons ignore surrounding whitespace and case. Without `group_key`, checking for overlap between external partitions is the caller's responsibility; DataFrame indices are not treated as row IDs.
105
+
71
106
  ### Numerical-stability recovery
72
107
 
73
108
  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.
@@ -15,7 +15,6 @@ from transformers import (
15
15
  AutoTokenizer,
16
16
  AutoConfig,
17
17
  AutoModelForSequenceClassification,
18
- TrainingArguments,
19
18
  DataCollatorWithPadding,
20
19
  EarlyStoppingCallback,
21
20
  set_seed,
@@ -38,6 +37,7 @@ from sklearn.preprocessing import label_binarize
38
37
  from mlflow.tracking import MlflowClient
39
38
 
40
39
  from bertuner.CustomTrainer import CustomTrainer
40
+ from bertuner.compat import TrainingArguments
41
41
  from bertuner.exceptions import NonFiniteTrainingError, NoStableTrialError
42
42
  from bertuner.TensorBoardCallback import (
43
43
  TensorBoardSyncCallback,
@@ -56,6 +56,7 @@ from bertuner.constants import (
56
56
  MODEL_DROPOUT_ATTRS,
57
57
  SEED,
58
58
  )
59
+ import inspect
59
60
 
60
61
 
61
62
  class BERTuneClassifier:
@@ -87,11 +88,10 @@ class BERTuneClassifier:
87
88
  retry_nonfinite_in_fp32: bool = True,
88
89
  max_grad_norm: float | None = 1.0,
89
90
  class_weight_warning_threshold: float | None = 100.0,
91
+ data_splits: dict[str, pd.DataFrame | str | os.PathLike] = None,
90
92
  ):
91
- if data_path is None and dataframe is None:
92
- raise ValueError("Provide either data_path (CSV) or dataframe, not neither.")
93
- if data_path is not None and dataframe is not None:
94
- raise ValueError("Provide either data_path (CSV) or dataframe, not both.")
93
+ if sum(source is not None for source in (data_path, dataframe, data_splits)) != 1:
94
+ raise ValueError("Provide exactly one of data_path (CSV), dataframe, or data_splits.")
95
95
  if precision not in {"auto", "fp32", "bf16", "fp16"}:
96
96
  raise ValueError("precision must be one of: 'auto', 'fp32', 'bf16', 'fp16'.")
97
97
  if max_grad_norm is not None and (
@@ -111,7 +111,12 @@ class BERTuneClassifier:
111
111
  self.text_feature = text_feature
112
112
  self.target_cols = target_cols
113
113
  self.seed = seed
114
- self.df = pd.read_csv(data_path) if data_path is not None else dataframe.copy()
114
+ self._data_splits = None
115
+ if data_splits is not None:
116
+ self._data_splits = self._load_data_splits(data_splits, group_key)
117
+ self.df = pd.concat(self._data_splits.values(), ignore_index=True)
118
+ else:
119
+ self.df = pd.read_csv(data_path) if data_path is not None else dataframe.copy()
115
120
  # Missing/NaN text values are coerced to empty strings so tokenization
116
121
  # does not error on non-string inputs. Warn so callers know rows were
117
122
  # altered rather than dropped.
@@ -123,6 +128,9 @@ class BERTuneClassifier:
123
128
  stacklevel=2,
124
129
  )
125
130
  self.df[text_feature] = self.df[text_feature].fillna("").astype(str)
131
+ if self._data_splits is not None:
132
+ for frame in self._data_splits.values():
133
+ frame[text_feature] = frame[text_feature].fillna("").astype(str)
126
134
  if num_labels is not None:
127
135
  self.num_labels = num_labels
128
136
  elif self.is_multilabel:
@@ -163,7 +171,7 @@ class BERTuneClassifier:
163
171
  else:
164
172
  print(f"Logging mlflow runs locally to: {self.mlflow_uri} (no server needed)")
165
173
 
166
- if self.is_multilabel and self.group_key:
174
+ if self.is_multilabel and self.group_key and self._data_splits is None:
167
175
  print(
168
176
  "[WARNING] group_key is set but multi-label mode is active. "
169
177
  "StratifiedGroupKFold does not support multi-label targets — "
@@ -473,12 +481,51 @@ class BERTuneClassifier:
473
481
  # Data preparation
474
482
  # ------------------------------------------------------------------
475
483
 
484
+ def _load_data_splits(self, data_splits, group_key):
485
+ """Copy external partitions, validating their schema and group isolation."""
486
+ if not isinstance(data_splits, dict):
487
+ raise ValueError("data_splits must be a dict with train, validation, and test keys.")
488
+ sources = dict(data_splits)
489
+ if "val" in sources and "validation" not in sources:
490
+ sources["validation"] = sources.pop("val")
491
+ if set(sources) != {"train", "validation", "test"}:
492
+ raise ValueError("data_splits requires exactly train, validation (or val), and test.")
493
+ required = [self.text_feature, *self.target_cols]
494
+ if group_key is not None:
495
+ required.append(group_key)
496
+ frames = {}
497
+ for name in ("train", "validation", "test"):
498
+ source = sources[name]
499
+ if isinstance(source, pd.DataFrame):
500
+ frame = source.copy()
501
+ elif isinstance(source, (str, os.PathLike)):
502
+ frame = pd.read_csv(source)
503
+ else:
504
+ raise ValueError(f"data_splits['{name}'] must be a DataFrame or CSV path.")
505
+ missing = [col for col in required if col not in frame.columns]
506
+ if missing:
507
+ raise ValueError(f"Split '{name}' is missing required columns: {missing}")
508
+ if frame.empty:
509
+ raise ValueError(f"Split '{name}' must not be empty.")
510
+ frames[name] = frame
511
+ if group_key is not None:
512
+ seen = set()
513
+ for name, frame in frames.items():
514
+ if frame[group_key].isna().any():
515
+ raise ValueError(f"Split '{name}' contains missing group IDs in '{group_key}'.")
516
+ groups = set(frame[group_key].astype(str).str.strip().str.casefold())
517
+ if seen & groups:
518
+ raise ValueError(f"Group leakage detected in supplied splits for '{group_key}'.")
519
+ seen.update(groups)
520
+ return frames
521
+
476
522
  def _prepare_datasets(self, tokenizer, group_key, max_length=512):
477
523
  """
478
524
  Splits, balances, and tokenizes data.
479
525
 
480
526
  Splitting strategy
481
527
  ------------------
528
+ data_splits supplied → preserve external partitions and row order
482
529
  Multi-label + group_key → standard split (group stratification not supported
483
530
  for multi-label targets; warning shown in __init__)
484
531
  Single-label + group_key → StratifiedGroupKFold split
@@ -489,7 +536,11 @@ class BERTuneClassifier:
489
536
  not self.is_multilabel and group_key is not None and group_key in self.df.columns
490
537
  )
491
538
 
492
- if use_group_split:
539
+ if self._data_splits is not None:
540
+ train, val, test = (
541
+ self._data_splits[name] for name in ("train", "validation", "test")
542
+ )
543
+ elif use_group_split:
493
544
  self.df[group_key] = self.df[group_key].astype(str).str.strip().str.casefold()
494
545
  train, val, test = split_group_stratified(
495
546
  self.df,
@@ -696,7 +747,6 @@ class BERTuneClassifier:
696
747
  "gradient_checkpointing": self._use_gradient_checkpointing(max_length),
697
748
  "gradient_checkpointing_kwargs": {"use_reentrant": False},
698
749
  "weight_decay": params["weight_decay"],
699
- "warmup_ratio": params["warmup_ratio"],
700
750
  "metric_for_best_model": f"eval_{self.optimize_metric}",
701
751
  "greater_is_better": self.greater_is_better,
702
752
  "eval_strategy": "epoch",
@@ -710,17 +760,20 @@ class BERTuneClassifier:
710
760
  "seed": self.seed,
711
761
  **self._precision_flags(precision),
712
762
  }
713
- if final:
714
- # Canonical final metrics are logged explicitly after restoring the
715
- # best checkpoint; Trainer only sends loss curves to TensorBoard.
716
- kwargs.update(report_to=["tensorboard"], logging_dir=logging_dir)
717
- else:
718
- kwargs.update(
719
- lr_scheduler_type=params["scheduler"],
720
- remove_unused_columns=True,
721
- report_to=["none"],
722
- )
723
- return TrainingArguments(**kwargs)
763
+
764
+ _TA_PARAMS = inspect.signature(TrainingArguments.__init__).parameters
765
+
766
+ # v5 removed warmup_ratio and accepts fractions in warmup_steps.
767
+ warmup_key = "warmup_ratio" if "warmup_ratio" in _TA_PARAMS else "warmup_steps"
768
+ kwargs[warmup_key] = params["warmup_ratio"]
769
+ kwargs["lr_scheduler_type"] = params["scheduler"]
770
+ # Final training uses an explicit TensorBoard writer in _build_trainer;
771
+ # newer Transformers versions removed TrainingArguments.logging_dir.
772
+ kwargs.update(remove_unused_columns=True, report_to=["none"])
773
+ args = TrainingArguments(**kwargs)
774
+ if warmup_key == "warmup_steps":
775
+ args._bertuner_warmup_ratio = params["warmup_ratio"]
776
+ return args
724
777
 
725
778
  def _build_trainer(
726
779
  self,
@@ -742,8 +795,12 @@ class BERTuneClassifier:
742
795
  )
743
796
  ]
744
797
  if final:
798
+ from transformers.integrations import TensorBoardCallback
799
+ from torch.utils.tensorboard import SummaryWriter
800
+
745
801
  callbacks.extend(
746
802
  [
803
+ TensorBoardCallback(tb_writer=SummaryWriter(logging_dir)),
747
804
  TensorBoardSyncCallback(logging_dir),
748
805
  CleanupCheckpointsCallback,
749
806
  ]
@@ -807,11 +864,15 @@ class BERTuneClassifier:
807
864
  try:
808
865
  trainer.train()
809
866
  except Exception:
867
+ from transformers.integrations import TensorBoardCallback
868
+
810
869
  # Trainer does not emit on_train_end after an exception. Close only
811
870
  # BERTuner-owned writers before the caller decides whether to retry.
812
871
  for callback in trainer.callback_handler.callbacks:
813
872
  if isinstance(callback, TensorBoardSyncCallback):
814
873
  callback.writer.close()
874
+ elif isinstance(callback, TensorBoardCallback) and callback.tb_writer is not None:
875
+ callback.tb_writer.close()
815
876
  raise
816
877
  return trainer, model
817
878
 
@@ -69,6 +69,9 @@ class CustomTrainer(Trainer):
69
69
  callbacks = list(kwargs.pop("callbacks", None) or [])
70
70
  callbacks.append(NonFiniteGradientCallback(training_precision))
71
71
  super().__init__(callbacks=callbacks, **kwargs)
72
+ # compute_loss returns a microbatch mean and does not consume
73
+ # num_items_in_batch, even when the model's forward accepts **kwargs.
74
+ self.model_accepts_loss_kwargs = False
72
75
  self.loss_type = loss_type
73
76
  self.class_weights = class_weights
74
77
  self.training_precision = training_precision
@@ -1,6 +1,6 @@
1
1
  """bertuner: hyperparameter optimization and fine-tuning for BERT-style text classifiers."""
2
2
 
3
- __version__ = "0.2.1"
3
+ __version__ = "0.2.3"
4
4
 
5
5
  from bertuner.BERTuner import BERTuneClassifier
6
6
  from bertuner.Predictor import BERTunePredictor
@@ -0,0 +1,15 @@
1
+ """Small adapters for Transformers APIs shared by supported 4.x and 5.x."""
2
+
3
+ import math
4
+
5
+ from transformers import TrainingArguments as HFTrainingArguments
6
+
7
+
8
+ class TrainingArguments(HFTrainingArguments):
9
+ def get_warmup_steps(self, num_training_steps):
10
+ ratio = getattr(self, "_bertuner_warmup_ratio", None)
11
+ if ratio is not None:
12
+ # New warmup_steps treats 1.0 as one step, whereas the old
13
+ # warmup_ratio treats it as the entire training run.
14
+ return math.ceil(num_training_steps * ratio)
15
+ return super().get_warmup_steps(num_training_steps)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.2.1
3
+ Version: 0.2.3
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
@@ -58,6 +58,20 @@ git clone https://github.com/elemets/bertuner && cd bertuner
58
58
  pip install -r requirements.txt
59
59
  ```
60
60
 
61
+ BERTuner supports Transformers 4.48+ and 5.x. Compatibility tests cover the
62
+ 4.48 baseline and the latest available 4.x and 5.x releases, including final
63
+ training, TensorBoard logging, warmup scheduling, and gradient accumulation.
64
+ Future releases are checked by a weekly CI run rather than assumed compatible.
65
+
66
+ When initializing a classifier from a base pretrained model, a Transformers 5
67
+ load report may list `lm_head.*` weights as `UNEXPECTED` and classifier weights
68
+ as `MISSING`. This is expected when replacing the language-model head with a
69
+ classification head; the new head is learned during fine-tuning. Unexpected
70
+ encoder weights or shape mismatches should still be investigated.
71
+
72
+ Run the test suite with `python -m pytest tests -q`. Training regression tests
73
+ use a tiny local checkpoint; cached-model predictor tests skip if unavailable.
74
+
61
75
  MLflow tracking works in two modes:
62
76
 
63
77
  ```bash
@@ -109,6 +123,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
109
123
 
110
124
  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.
111
125
 
126
+ ### Using existing splits
127
+
128
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
129
+
130
+ ```python
131
+ classifier = BERTuneClassifier(
132
+ models_dir="../models/",
133
+ text_feature="text_col",
134
+ target_cols=["label_col"],
135
+ data_splits={
136
+ "train": train_df,
137
+ "validation": validation_df, # "val" also accepted
138
+ "test": test_df,
139
+ },
140
+ )
141
+ ```
142
+
143
+ Each value can also be a CSV path, such as `"data/train.csv"` (including `pathlib.Path` objects). All three splits must be nonempty and contain the text and target columns. Use the same numeric label encoding across all splits. Both optimization and final training preserve these partitions and their row order; no random splitting or stratification is performed. Class weights are calculated from training data, and threshold selection uses validation data.
144
+
145
+ If you supply `group_key`, each split must contain that column with nonmissing IDs. Overlapping groups across splits raise an error, including in multi-label mode. Group comparisons ignore surrounding whitespace and case. Without `group_key`, checking for overlap between external partitions is the caller's responsibility; DataFrame indices are not treated as row IDs.
146
+
112
147
  ### Numerical-stability recovery
113
148
 
114
149
  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.
@@ -6,6 +6,7 @@ bertuner/CustomTrainer.py
6
6
  bertuner/Predictor.py
7
7
  bertuner/TensorBoardCallback.py
8
8
  bertuner/__init__.py
9
+ bertuner/compat.py
9
10
  bertuner/constants.py
10
11
  bertuner/exceptions.py
11
12
  bertuner/utils.py
@@ -17,4 +18,5 @@ bertuner.egg-info/top_level.txt
17
18
  tests/test_bertuner.py
18
19
  tests/test_numerical_stability.py
19
20
  tests/test_predictor.py
21
+ tests/test_transformers_compat.py
20
22
  tests/test_utils.py
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "bertuner"
7
- version = "0.2.1"
7
+ version = "0.2.3"
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" }
@@ -85,6 +85,93 @@ class TestConstruction:
85
85
  assert clf.num_labels == 7
86
86
 
87
87
 
88
+ class TestExternalSplits:
89
+ @staticmethod
90
+ def splits():
91
+ df = make_df(n=12)
92
+ return {"train": df.iloc[:6], "validation": df.iloc[6:9], "test": df.iloc[9:]}
93
+
94
+ @pytest.mark.parametrize("multilabel", [False, True])
95
+ @pytest.mark.parametrize("csv", [False, True])
96
+ def test_preserves_partitions_across_preparations(self, tmp_path, monkeypatch, multilabel, csv):
97
+ splits = self.splits()
98
+ targets = ["target"]
99
+ if multilabel:
100
+ targets.append("other")
101
+ splits = {name: frame.assign(other=1 - frame.target) for name, frame in splits.items()}
102
+ sources = splits
103
+ if csv:
104
+ sources = {}
105
+ for name, frame in splits.items():
106
+ sources[name] = tmp_path / f"{name}.csv"
107
+ frame.to_csv(sources[name], index=False)
108
+ clf = make_classifier(tmp_path, dataframe=None, data_splits=sources, target_cols=targets)
109
+ def no_split(*args, **kwargs):
110
+ pytest.fail("External partitions must never be split again")
111
+ monkeypatch.setattr("bertuner.BERTuner.train_val_test_split", no_split)
112
+ monkeypatch.setattr("bertuner.BERTuner.split_group_stratified", no_split)
113
+ def tokenizer(texts, **kwargs):
114
+ return {"input_ids": [[int(text.rsplit(" ", 1)[1])] for text in texts]}
115
+ for seed in (7, 99):
116
+ clf.seed = seed
117
+ datasets = clf._prepare_datasets(tokenizer, None)
118
+ for ds, frame in zip(datasets, splits.values()):
119
+ assert ds[:]["input_ids"].tolist() == [[i] for i in frame.index]
120
+ expected = frame[targets].values.tolist() if multilabel else frame.target.tolist()
121
+ assert ds[:]["labels"].tolist() == expected
122
+
123
+ def test_copies_inputs_and_cleans_text(self, tmp_path):
124
+ splits = self.splits()
125
+ splits["train"] = splits["train"].copy()
126
+ splits["train"].iloc[0, 0] = None
127
+ splits["val"] = splits.pop("validation")
128
+ with pytest.warns(UserWarning, match="missing"):
129
+ clf = make_classifier(tmp_path, dataframe=None, data_splits=splits)
130
+ assert clf._data_splits["train"].iloc[0, 0] == ""
131
+ assert splits["train"].iloc[0, 0] is None
132
+ splits["train"].iloc[1, 0] = "changed"
133
+ assert clf._data_splits["train"].iloc[1, 0] == "sample text 1"
134
+
135
+ @pytest.mark.parametrize("problem,match", [
136
+ ("missing_split", "requires exactly"),
137
+ ("extra_split", "requires exactly"),
138
+ ("empty", "must not be empty"),
139
+ ("column", "missing required columns"),
140
+ ("type", "DataFrame or CSV path"),
141
+ ])
142
+ def test_invalid_splits(self, tmp_path, problem, match):
143
+ splits = self.splits()
144
+ if problem == "missing_split":
145
+ splits.pop("test")
146
+ elif problem == "extra_split":
147
+ splits["val"] = splits["validation"]
148
+ elif problem == "empty":
149
+ splits["test"] = splits["test"].iloc[:0]
150
+ elif problem == "column":
151
+ splits["test"] = splits["test"].drop(columns="target")
152
+ else:
153
+ splits["test"] = None
154
+ with pytest.raises(ValueError, match=match):
155
+ make_classifier(tmp_path, dataframe=None, data_splits=splits)
156
+
157
+ def test_rejects_conflicting_sources(self, tmp_path):
158
+ with pytest.raises(ValueError, match="exactly one"):
159
+ make_classifier(tmp_path, data_splits=self.splits())
160
+
161
+ @pytest.mark.parametrize("multilabel", [False, True])
162
+ def test_group_validation(self, tmp_path, multilabel):
163
+ splits = {name: frame.assign(patient=name, other=1) for name, frame in self.splits().items()}
164
+ kwargs = dict(dataframe=None, data_splits=splits, group_key="patient",
165
+ target_cols=["target", "other"] if multilabel else ["target"])
166
+ make_classifier(tmp_path, **kwargs)
167
+ splits["test"]["patient"] = " TRAIN "
168
+ with pytest.raises(ValueError, match="Group leakage"):
169
+ make_classifier(tmp_path, **kwargs)
170
+ splits["test"]["patient"] = None
171
+ with pytest.raises(ValueError, match="missing group IDs"):
172
+ make_classifier(tmp_path, **kwargs)
173
+
174
+
88
175
  class TestMlflowUri:
89
176
  def test_plain_path_becomes_file_uri(self, tmp_path):
90
177
  clf = make_classifier(tmp_path, mlflow_tracking_uri="./mlruns")
@@ -396,6 +483,7 @@ class TestLossCurveLogging:
396
483
  ]
397
484
  )
398
485
 
486
+ (tmp_path / "downloaded-artifacts").mkdir()
399
487
  artifact_path = MlflowClient().download_artifacts(
400
488
  run.info.run_id,
401
489
  "plots/training_vs_evaluation_loss.png",
@@ -698,7 +786,7 @@ class TestMulticlassWeightedLoss:
698
786
  loss = trainer._singlelabel_loss(logits, labels, torch.device("cpu"))
699
787
  assert torch.isfinite(loss)
700
788
 
701
- def test_end_to_end_training_three_classes(self, tmp_path):
789
+ def test_end_to_end_training_three_classes(self, tmp_path, tiny_model_path):
702
790
  import torch
703
791
  from transformers import (
704
792
  AutoModelForSequenceClassification,
@@ -708,7 +796,7 @@ class TestMulticlassWeightedLoss:
708
796
  )
709
797
  from bertuner.CustomTrainer import CustomTrainer
710
798
 
711
- model_path = "prajjwal1/bert-tiny"
799
+ model_path = tiny_model_path
712
800
  clf = make_classifier(tmp_path, dataframe=make_df(n=60, num_classes=3))
713
801
  assert clf.num_labels == 3
714
802
 
@@ -0,0 +1,88 @@
1
+ """Run unchanged against the oldest supported Transformers and current releases."""
2
+ import copy
3
+ from types import SimpleNamespace
4
+
5
+ import numpy as np
6
+ import pytest
7
+ import torch
8
+ from datasets import Dataset
9
+ from tensorboard.backend.event_processing.event_accumulator import EventAccumulator
10
+ from transformers import AutoTokenizer, TrainingArguments
11
+
12
+ from bertuner.CustomTrainer import CustomTrainer
13
+ from test_bertuner import make_classifier, make_df
14
+ from test_numerical_stability import sampled_params
15
+
16
+
17
+ @pytest.mark.parametrize("ratio", [0.0, 0.1, 1.0])
18
+ @pytest.mark.parametrize("final", [False, True])
19
+ def test_warmup_and_scheduler(tmp_path, ratio, final):
20
+ clf = make_classifier(tmp_path)
21
+ clf.optimize_metric = "avg_precision"
22
+ clf.greater_is_better = True
23
+ params = dict(sampled_params(), warmup_ratio=ratio, scheduler="cosine")
24
+ args = clf._build_training_arguments(
25
+ params, str(tmp_path / "output"), 32, "fp32", final=final,
26
+ logging_dir=str(tmp_path / "logs"),
27
+ )
28
+ assert args.get_warmup_steps(100) == int(100 * ratio)
29
+ assert args.lr_scheduler_type == "cosine"
30
+
31
+
32
+ def test_final_training_logs_and_reloads(tmp_path, tiny_model_path):
33
+ clf = make_classifier(tmp_path, dataframe=make_df(n=30, num_classes=3))
34
+ clf.optimize_metric = "f1"
35
+ clf.greater_is_better = True
36
+ tokenizer = AutoTokenizer.from_pretrained(tiny_model_path)
37
+ train, val, _ = clf._prepare_datasets(tokenizer, None, max_length=32)
38
+ params = sampled_params()
39
+ log_dir = str(tmp_path / "logs")
40
+ args = clf._build_training_arguments(
41
+ params, str(tmp_path / "output"), 32, "fp32", final=True, logging_dir=log_dir,
42
+ )
43
+ args.num_train_epochs = 1
44
+ args.use_cpu = True
45
+ trainer = clf._build_trainer(
46
+ clf._load_model(tiny_model_path, 0.0), args, train, val, tokenizer,
47
+ params, clf._compute_class_weights(train), "fp32", final=True,
48
+ logging_dir=log_dir,
49
+ )
50
+ assert np.isfinite(trainer.train().training_loss)
51
+ assert trainer.state.best_model_checkpoint is not None
52
+ assert np.isfinite(trainer.evaluate()["eval_loss"])
53
+ events = EventAccumulator(log_dir).Reload()
54
+ assert "train/loss" in events.Tags()["scalars"]
55
+ assert "eval/loss" in events.Tags()["scalars"]
56
+
57
+
58
+ class KwargsClassifier(torch.nn.Module):
59
+ def __init__(self):
60
+ super().__init__()
61
+ self.linear = torch.nn.Linear(2, 2)
62
+ self.accepts_loss_kwargs = True
63
+
64
+ def forward(self, features, labels=None, **kwargs):
65
+ return SimpleNamespace(logits=self.linear(features))
66
+
67
+
68
+ def test_accumulation_matches_full_batch_update(tmp_path):
69
+ torch.manual_seed(42)
70
+ initial = KwargsClassifier()
71
+ dataset = Dataset.from_dict({"features": [[1., 2.]] * 4, "labels": [1] * 4})
72
+ models = []
73
+ for batch, accumulation in [(4, 1), (2, 2)]:
74
+ model = copy.deepcopy(initial)
75
+ args = TrainingArguments(
76
+ output_dir=str(tmp_path / str(batch)), use_cpu=True,
77
+ per_device_train_batch_size=batch, gradient_accumulation_steps=accumulation,
78
+ max_steps=1, max_grad_norm=0., report_to="none", save_strategy="no",
79
+ lr_scheduler_type="constant", disable_tqdm=True,
80
+ )
81
+ trainer = CustomTrainer(
82
+ model=model, args=args, train_dataset=dataset, loss_type="plain",
83
+ optimizers=(torch.optim.SGD(model.parameters(), lr=0.1), None),
84
+ )
85
+ trainer.train()
86
+ models.append(model)
87
+ for full, accumulated in zip(models[0].parameters(), models[1].parameters()):
88
+ torch.testing.assert_close(full, accumulated)
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes