fixes
This commit is contained in:
@@ -863,7 +863,8 @@ class Flux(nn.Module):
|
||||
|
||||
self.time_in.enable_gradient_checkpointing()
|
||||
self.vector_in.enable_gradient_checkpointing()
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks + self.single_blocks:
|
||||
block.enable_gradient_checkpointing()
|
||||
@@ -875,7 +876,8 @@ class Flux(nn.Module):
|
||||
|
||||
self.time_in.disable_gradient_checkpointing()
|
||||
self.vector_in.disable_gradient_checkpointing()
|
||||
self.guidance_in.disable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks + self.single_blocks:
|
||||
block.disable_gradient_checkpointing()
|
||||
@@ -972,7 +974,8 @@ class FluxUpper(nn.Module):
|
||||
|
||||
self.time_in.enable_gradient_checkpointing()
|
||||
self.vector_in.enable_gradient_checkpointing()
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks:
|
||||
block.enable_gradient_checkpointing()
|
||||
@@ -984,7 +987,8 @@ class FluxUpper(nn.Module):
|
||||
|
||||
self.time_in.disable_gradient_checkpointing()
|
||||
self.vector_in.disable_gradient_checkpointing()
|
||||
self.guidance_in.disable_gradient_checkpointing()
|
||||
if self.guidance_in.__class__ != nn.Identity:
|
||||
self.guidance_in.enable_gradient_checkpointing()
|
||||
|
||||
for block in self.double_blocks:
|
||||
block.disable_gradient_checkpointing()
|
||||
|
||||
@@ -28,6 +28,7 @@ logger = logging.getLogger(__name__)
|
||||
class FluxNetworkTrainer(NetworkTrainer):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.sample_prompts_te_outputs = None
|
||||
|
||||
def assert_extra_args(self, args, train_dataset_group):
|
||||
super().assert_extra_args(args, train_dataset_group)
|
||||
@@ -41,11 +42,17 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
args.network_train_unet_only or not args.cache_text_encoder_outputs
|
||||
), "network for Text Encoder cannot be trained with caching Text Encoder outputs"
|
||||
|
||||
if args.max_token_length is not None:
|
||||
logger.warning("max_token_length is not used in Flux training")
|
||||
|
||||
train_dataset_group.verify_bucket_reso_steps(32) # TODO check this
|
||||
|
||||
def get_flux_model_name(self, args):
|
||||
return "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev"
|
||||
|
||||
def load_target_model(self, args, weight_dtype, accelerator):
|
||||
# currently offload to cpu for some models
|
||||
name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" # TODO change this to a more robust way
|
||||
name = self.get_flux_model_name(args)
|
||||
# if we load to cpu, flux.to(fp8) takes a long time
|
||||
model = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu")
|
||||
|
||||
@@ -101,7 +108,18 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
return flux_lower
|
||||
|
||||
def get_tokenize_strategy(self, args):
|
||||
return strategy_flux.FluxTokenizeStrategy(args.max_token_length, args.tokenizer_cache_dir)
|
||||
name = self.get_flux_model_name(args)
|
||||
|
||||
if args.t5xxl_max_token_length is None:
|
||||
if name == "schnell":
|
||||
t5xxl_max_token_length = 256
|
||||
else:
|
||||
t5xxl_max_token_length = 512
|
||||
else:
|
||||
t5xxl_max_token_length = args.t5xxl_max_token_length
|
||||
|
||||
logger.info(f"t5xxl_max_token_length: {t5xxl_max_token_length}")
|
||||
return strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length, args.tokenizer_cache_dir)
|
||||
|
||||
def get_tokenizers(self, tokenize_strategy: strategy_flux.FluxTokenizeStrategy):
|
||||
return [tokenize_strategy.clip_l, tokenize_strategy.t5xxl]
|
||||
@@ -145,7 +163,7 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process)
|
||||
|
||||
# cache sample prompts
|
||||
self.sample_prompts_te_outputs = None
|
||||
|
||||
if args.sample_prompts is not None:
|
||||
logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}")
|
||||
|
||||
@@ -549,6 +567,7 @@ class InitFluxTraining:
|
||||
"network_train_unet_only": ("BOOLEAN", {"default": True, "tooltip": "wheter to train the text encoder"}),
|
||||
"text_encoder_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}),
|
||||
"apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "apply t5 attention mask"}),
|
||||
"t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev uses 512, schnell 256"}),
|
||||
"cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}),
|
||||
"split_mode": ("BOOLEAN", {"default": False, "tooltip": "[EXPERIMENTAL] use split mode for Flux model, network arg `train_blocks=single` is required"}),
|
||||
@@ -561,7 +580,7 @@ class InitFluxTraining:
|
||||
"model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)."}),
|
||||
"discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}),
|
||||
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
|
||||
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "default", "tooltip": "memory efficient attention mode"}),
|
||||
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}),
|
||||
"sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
|
||||
},
|
||||
}
|
||||
@@ -685,27 +704,11 @@ class FluxTrainLoop:
|
||||
target_global_step = network_trainer.global_step + steps
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
while network_trainer.global_step < target_global_step:
|
||||
epoch = training_loop(
|
||||
break_at_steps=target_global_step,
|
||||
epoch=network_trainer.current_epoch.value,
|
||||
num_train_epochs=network_trainer.num_train_epochs,
|
||||
accelerator=network_trainer.accelerator,
|
||||
network=network_trainer.network,
|
||||
text_encoder=network_trainer.text_encoder,
|
||||
unet=network_trainer.unet,
|
||||
vae=network_trainer.vae,
|
||||
tokenizers=network_trainer.tokenizers,
|
||||
args=network_trainer.args,
|
||||
train_dataloader=network_trainer.train_dataloader,
|
||||
initial_step=network_trainer.initial_step,
|
||||
global_step=network_trainer.global_step,
|
||||
current_epoch=network_trainer.current_epoch,
|
||||
metadata=network_trainer.metadata,
|
||||
optimizer=network_trainer.optimizer,
|
||||
lr_scheduler=network_trainer.lr_scheduler,
|
||||
loss_recorder=network_trainer.loss_recorder
|
||||
steps_done = training_loop(
|
||||
break_at_steps = target_global_step,
|
||||
epoch = network_trainer.current_epoch.value,
|
||||
)
|
||||
pbar.update(network_trainer.global_step - initial_global_step)
|
||||
pbar.update(steps_done)
|
||||
|
||||
# Also break if the global steps have reached the max train steps
|
||||
if network_trainer.global_step >= network_trainer.args.max_train_steps:
|
||||
|
||||
+34
-33
@@ -1025,13 +1025,30 @@ class NetworkTrainer:
|
||||
|
||||
clean_memory_on_device(accelerator.device)
|
||||
|
||||
self.epoch_to_start = epoch_to_start
|
||||
self.num_train_epochs = num_train_epochs
|
||||
self.accelerator = accelerator
|
||||
self.network = network
|
||||
self.text_encoder = text_encoder
|
||||
self.unet = unet
|
||||
self.vae = vae
|
||||
self.tokenizers = tokenizers
|
||||
self.args = args
|
||||
self.train_dataloader = train_dataloader
|
||||
self.initial_step = initial_step
|
||||
self.current_epoch = current_epoch
|
||||
self.metadata = metadata
|
||||
self.optimizer = optimizer
|
||||
self.lr_scheduler = lr_scheduler
|
||||
self.loss_recorder = loss_recorder
|
||||
self.save_model = save_model
|
||||
self.remove_model = remove_model
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps"
|
||||
)
|
||||
def training_loop(break_at_steps, epoch, num_train_epochs, accelerator, network, text_encoder,
|
||||
unet, vae, tokenizers, args, train_dataloader, initial_step, global_step,
|
||||
current_epoch, metadata, optimizer, lr_scheduler, loss_recorder):
|
||||
|
||||
def training_loop(break_at_steps, epoch):
|
||||
steps_done = 0
|
||||
accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
|
||||
current_epoch.value = epoch + 1
|
||||
|
||||
@@ -1040,14 +1057,14 @@ class NetworkTrainer:
|
||||
accelerator.unwrap_model(network).on_epoch_start(text_encoder, unet)
|
||||
|
||||
skipped_dataloader = None
|
||||
if initial_step > 0:
|
||||
skipped_dataloader = accelerator.skip_first_batches(train_dataloader, initial_step - 1)
|
||||
initial_step = 1
|
||||
if self.initial_step > 0:
|
||||
skipped_dataloader = accelerator.skip_first_batches(train_dataloader, self.initial_step - 1)
|
||||
self.initial_step = 1
|
||||
|
||||
for step, batch in enumerate(skipped_dataloader or train_dataloader):
|
||||
current_step.value = global_step
|
||||
if initial_step > 0:
|
||||
initial_step -= 1
|
||||
current_step.value = self.global_step
|
||||
if self.initial_step > 0:
|
||||
self.initial_step -= 1
|
||||
continue
|
||||
|
||||
with accelerator.accumulate(training_model):
|
||||
@@ -1158,7 +1175,7 @@ class NetworkTrainer:
|
||||
# Checks if the accelerator has performed an optimization step behind the scenes
|
||||
if accelerator.sync_gradients:
|
||||
progress_bar.update(1)
|
||||
global_step += 1
|
||||
self.global_step += 1
|
||||
|
||||
current_loss = loss.detach().item()
|
||||
loss_recorder.add(epoch=epoch, step=step, loss=current_loss)
|
||||
@@ -1173,35 +1190,19 @@ class NetworkTrainer:
|
||||
logs = self.generate_step_logs(
|
||||
args, current_loss, avr_loss, lr_scheduler, lr_descriptions, keys_scaled, mean_norm, maximum_norm
|
||||
)
|
||||
accelerator.log(logs, step=global_step)
|
||||
accelerator.log(logs, step=self.global_step)
|
||||
|
||||
if global_step >= break_at_steps:
|
||||
if self.global_step >= break_at_steps:
|
||||
break
|
||||
steps_done += 1
|
||||
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss/epoch": loss_recorder.moving_average}
|
||||
accelerator.log(logs, step=epoch + 1)
|
||||
self.global_step = global_step
|
||||
return current_epoch.value
|
||||
|
||||
return steps_done
|
||||
|
||||
|
||||
self.epoch_to_start = epoch_to_start
|
||||
self.num_train_epochs = num_train_epochs
|
||||
self.accelerator = accelerator
|
||||
self.network = network
|
||||
self.text_encoder = text_encoder
|
||||
self.unet = unet
|
||||
self.vae = vae
|
||||
self.tokenizers = tokenizers
|
||||
self.args = args
|
||||
self.train_dataloader = train_dataloader
|
||||
self.initial_step = initial_step
|
||||
self.current_epoch = current_epoch
|
||||
self.metadata = metadata
|
||||
self.optimizer = optimizer
|
||||
self.lr_scheduler = lr_scheduler
|
||||
self.loss_recorder = loss_recorder
|
||||
self.save_model = save_model
|
||||
self.remove_model = remove_model
|
||||
|
||||
return training_loop
|
||||
|
||||
|
||||
Reference in New Issue
Block a user