diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py index 8512eef..64c9ddd 100644 --- a/toolkit/config_modules.py +++ b/toolkit/config_modules.py @@ -915,6 +915,7 @@ class DatasetConfig: """ def __init__(self, **kwargs): + self.batch_size: Union[int, None] = kwargs.get('batch_size', None) self.type = kwargs.get('type', 'image') # sd, slider, reference # will be legacy self.folder_path: str = kwargs.get('folder_path', None) @@ -1504,7 +1505,4 @@ def validate_configs( raise ValueError(f"Cannot cache unload text encoder with {model_config.arch} model. Control images are encoded with text embeddings. You can cache the text embeddings though") if train_config.diff_output_preservation and train_config.blank_prompt_preservation: - raise ValueError("Cannot use both differential output preservation and blank prompt preservation at the same time. Please set one of them to False.") - - if train_config.batch_size > 1 and any(dataset_config.auto_frame_count for dataset_config in dataset_configs): - raise ValueError("Cannot use batch size greater than 1 with auto_frame_count. Please set batch_size to 1 or auto_frame_count to False.") + raise ValueError("Cannot use both differential output preservation and blank prompt preservation at the same time. Please set one of them to False.") \ No newline at end of file diff --git a/toolkit/data_loader.py b/toolkit/data_loader.py index 3e9773a..7be76be 100644 --- a/toolkit/data_loader.py +++ b/toolkit/data_loader.py @@ -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, diff --git a/ui/src/app/jobs/new/SimpleJob.tsx b/ui/src/app/jobs/new/SimpleJob.tsx index d8c15e2..e82cd03 100644 --- a/ui/src/app/jobs/new/SimpleJob.tsx +++ b/ui/src/app/jobs/new/SimpleJob.tsx @@ -1232,6 +1232,17 @@ export default function SimpleJob({ placeholder="eg. 1" docKey={'dataset.num_repeats'} /> + + setJobConfig(value == null ? undefined : value, `config.process[0].datasets[${i}].batch_size`) + } + placeholder={`${jobConfig.config.process[0].train.batch_size}`} + min={1} + allowEmpty + />
void; min?: number; max?: number; + // when true, clearing the input calls onChange(null) instead of being ignored + allowEmpty?: boolean; } export const NumberInput = (props: NumberInputProps) => { - const { label, value, onChange, placeholder, required, min, max, docKey = null } = props; + const { label, value, onChange, placeholder, required, min, max, allowEmpty, docKey = null } = props; let { doc } = props; if (!doc && docKey) { doc = getDoc(docKey); @@ -193,6 +195,9 @@ export const NumberInput = (props: NumberInputProps) => { // Handle empty or partial inputs if (rawValue === '' || rawValue === '-') { // For empty or partial negative input, don't call onChange yet + if (rawValue === '' && allowEmpty) { + onChange(null); + } return; } diff --git a/ui/src/types.ts b/ui/src/types.ts index 53cca38..bcb27c0 100644 --- a/ui/src/types.ts +++ b/ui/src/types.ts @@ -106,6 +106,7 @@ export interface SaveConfig { } export interface DatasetConfig { + batch_size?: number; folder_path: string; mask_path: string | null; mask_min_value: number;