Skip to content

Commit 1a0c94f

Browse files
authored
Merge pull request #3089 from bghira/agent/dataset-batch-strict-validation
Validate dataset train batch sizes strictly
2 parents 69c6823 + b91906f commit 1a0c94f

4 files changed

Lines changed: 114 additions & 26 deletions

File tree

simpletuner/helpers/data_backend/config/image.py

Lines changed: 4 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from dataclasses import dataclass, field
55
from typing import Any, Dict, List, Optional, Union
66

7-
from simpletuner.helpers.data_backend.dataset_types import DatasetType
7+
from simpletuner.helpers.data_backend.dataset_types import DatasetType, parse_positive_train_batch_size
88
from simpletuner.helpers.training.state_tracker import StateTracker
99

1010
from . import validators
@@ -213,10 +213,7 @@ def _get_arg(key: str, default: Any = None) -> Any:
213213
config.disable_validation = backend_dict.get("disable_validation", False)
214214
train_batch_size = backend_dict.get("train_batch_size")
215215
if train_batch_size not in (None, ""):
216-
try:
217-
config.train_batch_size = int(train_batch_size)
218-
except (TypeError, ValueError) as exc:
219-
raise ValueError(f"(id={config.id}) train_batch_size must be a positive integer.") from exc
216+
config.train_batch_size = parse_positive_train_batch_size(train_batch_size, config.id)
220217
if "hash_filenames" in backend_dict and config.backend_type != "csv":
221218
config.hash_filenames = backend_dict.get("hash_filenames")
222219
config.source_dataset_id = backend_dict.get("source_dataset_id")
@@ -372,8 +369,8 @@ def validate(self, args: Dict[str, Any]) -> None:
372369
self._validate_video_settings(args)
373370

374371
validators.check_for_caption_filter_list_misuse(self.dataset_type, False, self.id)
375-
if self.train_batch_size is not None and self.train_batch_size < 1:
376-
raise ValueError(f"(id={self.id}) train_batch_size must be a positive integer.")
372+
if self.train_batch_size is not None:
373+
self.train_batch_size = parse_positive_train_batch_size(self.train_batch_size, self.id)
377374

378375
def _validate_controlnet_requirements(self, args: Dict[str, Any]) -> None:
379376
def _get_controlnet_flag(source: Any) -> Optional[bool]:

simpletuner/helpers/data_backend/dataset_types.py

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
from enum import Enum
6+
from numbers import Integral
67
from typing import Any, Iterable, Mapping, Optional, Sequence
78

89

@@ -72,6 +73,26 @@ def get_arg_value(args: Any, key: str, default: Any = None) -> Any:
7273
return getattr(args, key, default)
7374

7475

76+
def parse_positive_train_batch_size(raw_value: Any, backend_id: Optional[str] = None) -> int:
77+
"""Parse a positive integer dataset training batch size."""
78+
error_message = f"(id={backend_id}) train_batch_size must be a positive integer."
79+
80+
if isinstance(raw_value, bool):
81+
raise ValueError(error_message)
82+
if isinstance(raw_value, Integral):
83+
batch_size = int(raw_value)
84+
elif isinstance(raw_value, str) and raw_value.isascii() and raw_value.isdecimal():
85+
batch_size = int(raw_value)
86+
if raw_value != str(batch_size):
87+
raise ValueError(error_message)
88+
else:
89+
raise ValueError(error_message)
90+
91+
if batch_size < 1:
92+
raise ValueError(error_message)
93+
return batch_size
94+
95+
7596
def resolve_dataset_train_batch_size(
7697
backend: Mapping[str, Any],
7798
args: Any,
@@ -88,12 +109,4 @@ def resolve_dataset_train_batch_size(
88109
raw_value = get_arg_value(args, "train_batch_size", 1)
89110

90111
resolved_backend_id = backend_id if backend_id is not None else backend.get("id")
91-
try:
92-
batch_size = int(raw_value)
93-
except (TypeError, ValueError) as exc:
94-
raise ValueError(f"(id={resolved_backend_id}) train_batch_size must be a positive integer.") from exc
95-
96-
if batch_size < 1:
97-
raise ValueError(f"(id={resolved_backend_id}) train_batch_size must be a positive integer.")
98-
99-
return batch_size
112+
return parse_positive_train_batch_size(raw_value, resolved_backend_id)

tests/test_backend_config.py

Lines changed: 34 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,8 @@
88
from typing import Any, Dict
99
from unittest.mock import Mock, patch
1010

11+
import numpy as np
12+
1113
from simpletuner.helpers.data_backend.config import (
1214
BaseBackendConfig,
1315
ImageBackendConfig,
@@ -335,18 +337,40 @@ def test_train_batch_size_round_trip(self):
335337
self.assertEqual(config.train_batch_size, 3)
336338
self.assertEqual(output["config"]["train_batch_size"], 3)
337339

338-
def test_train_batch_size_must_be_positive(self):
339-
backend_dict = {
340-
"id": "image_test",
341-
"type": "local",
342-
"dataset_type": "image",
343-
"train_batch_size": 0,
344-
}
340+
def test_train_batch_size_must_be_a_positive_integer(self):
341+
invalid_values = (True, False, 2.5, 3.0, "2.5", 0, "0", -1, "-1", "03", "+3", "three")
342+
343+
for value in invalid_values:
344+
with self.subTest(value=value):
345+
backend_dict = {
346+
"id": "image_test",
347+
"type": "local",
348+
"dataset_type": "image",
349+
"train_batch_size": value,
350+
}
351+
352+
with self.assertRaisesRegex(
353+
ValueError,
354+
r"\(id=image_test\) train_batch_size must be a positive integer\.",
355+
):
356+
ImageBackendConfig.from_dict(backend_dict, self.args)
357+
358+
def test_validate_rejects_invalid_direct_train_batch_size(self):
359+
config = ImageBackendConfig(id="image_test", train_batch_size=True)
360+
361+
with self.assertRaisesRegex(
362+
ValueError,
363+
r"\(id=image_test\) train_batch_size must be a positive integer\.",
364+
):
365+
config.validate(self.args)
345366

346-
config = ImageBackendConfig.from_dict(backend_dict, self.args)
367+
def test_validate_normalizes_integral_train_batch_size(self):
368+
config = ImageBackendConfig(id="image_test", train_batch_size=np.int64(3))
347369

348-
with self.assertRaisesRegex(ValueError, "train_batch_size"):
349-
config.validate(self.args)
370+
config.validate(self.args)
371+
372+
self.assertIs(type(config.train_batch_size), int)
373+
self.assertIs(type(config.to_dict()["config"]["train_batch_size"]), int)
350374

351375
def test_vae_cache_ondemand_round_trip(self):
352376
backend_dict = {

tests/test_dataset_types.py

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
"""Tests for dataset type helpers."""
2+
3+
import unittest
4+
5+
from simpletuner.helpers.data_backend.dataset_types import (
6+
DatasetType,
7+
parse_positive_train_batch_size,
8+
resolve_dataset_train_batch_size,
9+
)
10+
11+
12+
class DatasetTrainBatchSizeTestCase(unittest.TestCase):
13+
def test_parser_accepts_positive_integers_and_canonical_strings(self):
14+
for value in (1, 3, "1", "3"):
15+
with self.subTest(value=value):
16+
self.assertEqual(parse_positive_train_batch_size(value, "dataset"), int(value))
17+
18+
def test_parser_rejects_non_positive_or_non_integer_values(self):
19+
invalid_values = (True, False, 2.5, 3.0, "2.5", 0, "0", -1, "-1", "03", "+3", "three", None)
20+
21+
for value in invalid_values:
22+
with self.subTest(value=value):
23+
with self.assertRaisesRegex(
24+
ValueError,
25+
r"\(id=dataset\) train_batch_size must be a positive integer\.",
26+
):
27+
parse_positive_train_batch_size(value, "dataset")
28+
29+
def test_resolver_validates_dataset_and_global_values(self):
30+
with self.assertRaisesRegex(ValueError, "train_batch_size"):
31+
resolve_dataset_train_batch_size(
32+
{"id": "dataset", "dataset_type": "image", "train_batch_size": 2.5},
33+
{"train_batch_size": 1},
34+
)
35+
36+
with self.assertRaisesRegex(ValueError, "train_batch_size"):
37+
resolve_dataset_train_batch_size(
38+
{"id": "dataset", "dataset_type": "image"},
39+
{"train_batch_size": True},
40+
)
41+
42+
def test_eval_resolver_forces_batch_size_one_before_validation(self):
43+
self.assertEqual(
44+
resolve_dataset_train_batch_size(
45+
{"id": "eval", "dataset_type": "eval", "train_batch_size": 2.5},
46+
{"train_batch_size": True},
47+
dataset_type=DatasetType.EVAL,
48+
),
49+
1,
50+
)
51+
52+
53+
if __name__ == "__main__":
54+
unittest.main()

0 commit comments

Comments
 (0)