This commit is contained in:
kijai
2024-08-20 16:34:46 +03:00
parent 8e2669e809
commit d07378a99d
4 changed files with 27 additions and 6 deletions
+9 -5
View File
@@ -280,7 +280,8 @@ class FluxTrainer:
flux.requires_grad_(True)
if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None:
is_swapping_blocks = args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None
if is_swapping_blocks:
# Swap blocks between CPU and GPU to reduce memory usage, in forward and backward passes.
# This idea is based on 2kpr's great work. Thank you!
logger.info(
@@ -433,8 +434,11 @@ class FluxTrainer:
training_models = [ds_model]
else:
# acceleratorがなんかよろしくやってくれるらしい
flux = accelerator.prepare(flux)
# accelerator does some magic
# if we doesn't swap blocks, we can move the model to device
flux = accelerator.prepare(flux, device_placement=[not is_swapping_blocks])
if is_swapping_blocks:
flux.move_to_device_except_swap_blocks(accelerator.device) # reduce peak memory usage
optimizer, train_dataloader, lr_scheduler = accelerator.prepare(optimizer, train_dataloader, lr_scheduler)
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
@@ -560,7 +564,7 @@ class FluxTrainer:
init_kwargs=init_kwargs,
)
if args.double_blocks_to_swap is not None or args.single_blocks_to_swap is not None:
if is_swapping_blocks:
flux.prepare_block_swap_before_forward()
# For --sample_at_first
@@ -625,7 +629,7 @@ class FluxTrainer:
# get noisy model input and timesteps
noisy_model_input, timesteps, sigmas = flux_train_utils.get_noisy_model_input_and_timesteps(
args, noise_scheduler, latents, noise, accelerator.device, weight_dtype
args, noise_scheduler_copy, latents, noise, accelerator.device, weight_dtype
)
# pack latents and get img_ids
+16
View File
@@ -952,6 +952,22 @@ class Flux(nn.Module):
self.double_blocks_to_swap = double_blocks
self.single_blocks_to_swap = single_blocks
def move_to_device_except_swap_blocks(self, device: torch.device):
# assume model is on cpu
if self.double_blocks_to_swap:
save_double_blocks = self.double_blocks
self.double_blocks = None
if self.single_blocks_to_swap:
save_single_blocks = self.single_blocks
self.single_blocks = None
self.to(device)
if self.double_blocks_to_swap:
self.double_blocks = save_double_blocks
if self.single_blocks_to_swap:
self.single_blocks = save_single_blocks
def prepare_block_swap_before_forward(self):
# move last n blocks to cpu: they are on cuda
if self.double_blocks_to_swap:
+1 -1
View File
@@ -185,7 +185,7 @@ class InitFluxLoRATraining:
"network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}),
"network_alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 256.0, "step": 0.01, "tooltip": "network alpha"}),
"learning_rate": ("FLOAT", {"default": 4e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "learning rate"}),
"unet_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "unet learning rate"}),
#"unet_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "unet learning rate"}),
#"max_train_epochs": ("INT", {"default": 4, "min": 1, "max": 1000, "step": 1, "tooltip": "max number of training epochs"}),
"max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 10000, "step": 1, "tooltip": "max number of training steps"}),
#"network_train_unet_only": ("BOOLEAN", {"default": True, "tooltip": "wheter to train the text encoder"}),
+1
View File
@@ -316,6 +316,7 @@ class NetworkTrainer:
collator = train_util.collator_class(current_epoch, current_step, ds_for_collator)
if args.debug_dataset:
train_dataset_group.set_current_strategies()
train_util.debug_dataset(train_dataset_group)
return
if len(train_dataset_group) == 0: