Allow overriding the batch size on a per dataset level

This commit is contained in:
Jaret Burkett
2026-08-21 12:56:04 -06:00
parent 89102f76dc
commit b96476a841
5 changed files with 30 additions and 6 deletions

View File

@@ -704,7 +704,9 @@ def get_dataloader_from_datasets(
for config in dataset_config_list:
if config.type == 'image':
dataset = AiToolkitDataset(config, batch_size=batch_size, sd=sd)
# dataset level batch_size overrides the train config batch_size when set
dataset_batch_size = config.batch_size if config.batch_size is not None else batch_size
dataset = AiToolkitDataset(config, batch_size=dataset_batch_size, sd=sd)
datasets.append(dataset)
if config.buckets:
has_buckets = True
@@ -753,6 +755,13 @@ def get_dataloader_from_datasets(
**dataloader_kwargs
)
else:
# without buckets the dataloader batches across all datasets at once,
# so a dataset level batch_size cannot apply
for config in dataset_config_list:
if config.batch_size is not None:
raise ValueError(
f"Dataset level batch_size requires buckets to be enabled. Dataset {config.folder_path or config.dataset_path} has buckets disabled."
)
data_loader = DataLoader(
concatenated_dataset,
batch_size=batch_size,