bertuner 0.2.0__tar.gz → 0.2.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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: bertuner
3
- Version: 0.2.0
3
+ Version: 0.2.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
@@ -109,6 +109,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
109
109
 
110
110
  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
111
 
112
+ ### Using existing splits
113
+
114
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
115
+
116
+ ```python
117
+ classifier = BERTuneClassifier(
118
+ models_dir="../models/",
119
+ text_feature="text_col",
120
+ target_cols=["label_col"],
121
+ data_splits={
122
+ "train": train_df,
123
+ "validation": validation_df, # "val" also accepted
124
+ "test": test_df,
125
+ },
126
+ )
127
+ ```
128
+
129
+ 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.
130
+
131
+ 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.
132
+
112
133
  ### Numerical-stability recovery
113
134
 
114
135
  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.
@@ -68,6 +68,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
68
68
 
69
69
  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
70
 
71
+ ### Using existing splits
72
+
73
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
74
+
75
+ ```python
76
+ classifier = BERTuneClassifier(
77
+ models_dir="../models/",
78
+ text_feature="text_col",
79
+ target_cols=["label_col"],
80
+ data_splits={
81
+ "train": train_df,
82
+ "validation": validation_df, # "val" also accepted
83
+ "test": test_df,
84
+ },
85
+ )
86
+ ```
87
+
88
+ 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.
89
+
90
+ 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.
91
+
71
92
  ### Numerical-stability recovery
72
93
 
73
94
  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.
@@ -87,11 +87,10 @@ class BERTuneClassifier:
87
87
  retry_nonfinite_in_fp32: bool = True,
88
88
  max_grad_norm: float | None = 1.0,
89
89
  class_weight_warning_threshold: float | None = 100.0,
90
+ data_splits: dict[str, pd.DataFrame | str | os.PathLike] = None,
90
91
  ):
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.")
92
+ if sum(source is not None for source in (data_path, dataframe, data_splits)) != 1:
93
+ raise ValueError("Provide exactly one of data_path (CSV), dataframe, or data_splits.")
95
94
  if precision not in {"auto", "fp32", "bf16", "fp16"}:
96
95
  raise ValueError("precision must be one of: 'auto', 'fp32', 'bf16', 'fp16'.")
97
96
  if max_grad_norm is not None and (
@@ -111,7 +110,26 @@ class BERTuneClassifier:
111
110
  self.text_feature = text_feature
112
111
  self.target_cols = target_cols
113
112
  self.seed = seed
114
- self.df = pd.read_csv(data_path) if data_path is not None else dataframe.copy()
113
+ self._data_splits = None
114
+ if data_splits is not None:
115
+ self._data_splits = self._load_data_splits(data_splits, group_key)
116
+ self.df = pd.concat(self._data_splits.values(), ignore_index=True)
117
+ else:
118
+ self.df = pd.read_csv(data_path) if data_path is not None else dataframe.copy()
119
+ # Missing/NaN text values are coerced to empty strings so tokenization
120
+ # does not error on non-string inputs. Warn so callers know rows were
121
+ # altered rather than dropped.
122
+ n_missing = int(self.df[text_feature].isna().sum())
123
+ if n_missing:
124
+ warnings.warn(
125
+ f"{n_missing} missing (NaN) value(s) in text column "
126
+ f"'{text_feature}' were replaced with empty strings.",
127
+ stacklevel=2,
128
+ )
129
+ self.df[text_feature] = self.df[text_feature].fillna("").astype(str)
130
+ if self._data_splits is not None:
131
+ for frame in self._data_splits.values():
132
+ frame[text_feature] = frame[text_feature].fillna("").astype(str)
115
133
  if num_labels is not None:
116
134
  self.num_labels = num_labels
117
135
  elif self.is_multilabel:
@@ -152,7 +170,7 @@ class BERTuneClassifier:
152
170
  else:
153
171
  print(f"Logging mlflow runs locally to: {self.mlflow_uri} (no server needed)")
154
172
 
155
- if self.is_multilabel and self.group_key:
173
+ if self.is_multilabel and self.group_key and self._data_splits is None:
156
174
  print(
157
175
  "[WARNING] group_key is set but multi-label mode is active. "
158
176
  "StratifiedGroupKFold does not support multi-label targets — "
@@ -462,12 +480,51 @@ class BERTuneClassifier:
462
480
  # Data preparation
463
481
  # ------------------------------------------------------------------
464
482
 
483
+ def _load_data_splits(self, data_splits, group_key):
484
+ """Copy external partitions, validating their schema and group isolation."""
485
+ if not isinstance(data_splits, dict):
486
+ raise ValueError("data_splits must be a dict with train, validation, and test keys.")
487
+ sources = dict(data_splits)
488
+ if "val" in sources and "validation" not in sources:
489
+ sources["validation"] = sources.pop("val")
490
+ if set(sources) != {"train", "validation", "test"}:
491
+ raise ValueError("data_splits requires exactly train, validation (or val), and test.")
492
+ required = [self.text_feature, *self.target_cols]
493
+ if group_key is not None:
494
+ required.append(group_key)
495
+ frames = {}
496
+ for name in ("train", "validation", "test"):
497
+ source = sources[name]
498
+ if isinstance(source, pd.DataFrame):
499
+ frame = source.copy()
500
+ elif isinstance(source, (str, os.PathLike)):
501
+ frame = pd.read_csv(source)
502
+ else:
503
+ raise ValueError(f"data_splits['{name}'] must be a DataFrame or CSV path.")
504
+ missing = [col for col in required if col not in frame.columns]
505
+ if missing:
506
+ raise ValueError(f"Split '{name}' is missing required columns: {missing}")
507
+ if frame.empty:
508
+ raise ValueError(f"Split '{name}' must not be empty.")
509
+ frames[name] = frame
510
+ if group_key is not None:
511
+ seen = set()
512
+ for name, frame in frames.items():
513
+ if frame[group_key].isna().any():
514
+ raise ValueError(f"Split '{name}' contains missing group IDs in '{group_key}'.")
515
+ groups = set(frame[group_key].astype(str).str.strip().str.casefold())
516
+ if seen & groups:
517
+ raise ValueError(f"Group leakage detected in supplied splits for '{group_key}'.")
518
+ seen.update(groups)
519
+ return frames
520
+
465
521
  def _prepare_datasets(self, tokenizer, group_key, max_length=512):
466
522
  """
467
523
  Splits, balances, and tokenizes data.
468
524
 
469
525
  Splitting strategy
470
526
  ------------------
527
+ data_splits supplied → preserve external partitions and row order
471
528
  Multi-label + group_key → standard split (group stratification not supported
472
529
  for multi-label targets; warning shown in __init__)
473
530
  Single-label + group_key → StratifiedGroupKFold split
@@ -478,7 +535,11 @@ class BERTuneClassifier:
478
535
  not self.is_multilabel and group_key is not None and group_key in self.df.columns
479
536
  )
480
537
 
481
- if use_group_split:
538
+ if self._data_splits is not None:
539
+ train, val, test = (
540
+ self._data_splits[name] for name in ("train", "validation", "test")
541
+ )
542
+ elif use_group_split:
482
543
  self.df[group_key] = self.df[group_key].astype(str).str.strip().str.casefold()
483
544
  train, val, test = split_group_stratified(
484
545
  self.df,
@@ -1,6 +1,6 @@
1
1
  """bertuner: hyperparameter optimization and fine-tuning for BERT-style text classifiers."""
2
2
 
3
- __version__ = "0.2.0"
3
+ __version__ = "0.2.2"
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.0
3
+ Version: 0.2.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
@@ -109,6 +109,27 @@ Multi-label classification: pass several target columns — `target_cols=["l1",
109
109
 
110
110
  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
111
 
112
+ ### Using existing splits
113
+
114
+ Supply `data_splits` instead of `data_path` or `dataframe` to reuse your own partitions:
115
+
116
+ ```python
117
+ classifier = BERTuneClassifier(
118
+ models_dir="../models/",
119
+ text_feature="text_col",
120
+ target_cols=["label_col"],
121
+ data_splits={
122
+ "train": train_df,
123
+ "validation": validation_df, # "val" also accepted
124
+ "test": test_df,
125
+ },
126
+ )
127
+ ```
128
+
129
+ 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.
130
+
131
+ 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.
132
+
112
133
  ### Numerical-stability recovery
113
134
 
114
135
  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.
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "bertuner"
7
- version = "0.2.0"
7
+ version = "0.2.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" }
@@ -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")
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes