diff --git a/fine_tune.py b/fine_tune.py index 7ea76bc..ae7d9c8 100644 --- a/fine_tune.py +++ b/fine_tune.py @@ -175,7 +175,7 @@ def train(args): vae.requires_grad_(False) vae.eval() - train_dataset_group.new_cache_latents(vae, accelerator.is_main_process) + train_dataset_group.new_cache_latents(vae, accelerator) vae.to("cpu") clean_memory_on_device(accelerator.device) diff --git a/flux_train_comfy.py b/flux_train_comfy.py index 1412fba..50e10de 100644 --- a/flux_train_comfy.py +++ b/flux_train_comfy.py @@ -181,7 +181,7 @@ class FluxTrainer: ae.requires_grad_(False) ae.eval() - train_dataset_group.new_cache_latents(ae, accelerator.is_main_process) + train_dataset_group.new_cache_latents(ae, accelerator) ae.to("cpu") # if no sampling, vae can be deleted clean_memory_on_device(accelerator.device) diff --git a/library/flux_models.py b/library/flux_models.py index 46c1381..5871690 100644 --- a/library/flux_models.py +++ b/library/flux_models.py @@ -438,11 +438,11 @@ configs = { # region math - +from sageattention import sageattn def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, attn_mask: Optional[Tensor] = None) -> Tensor: q, k = apply_rope(q, k, pe) - x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask) + x = sageattn(q, k, v, attn_mask=attn_mask, smooth_k=False) x = rearrange(x, "B H L D -> B L (H D)") return x diff --git a/library/train_util.py b/library/train_util.py index dbb9ba8..24d7e10 100644 --- a/library/train_util.py +++ b/library/train_util.py @@ -1507,7 +1507,6 @@ class BaseDataset(torch.utils.data.Dataset): text_encoder_outputs = self.text_encoder_output_caching_strategy.load_outputs_npz( image_info.text_encoder_outputs_npz ) - text_encoder_outputs = [torch.FloatTensor(x) for x in text_encoder_outputs] else: tokenization_required = True text_encoder_outputs_list.append(text_encoder_outputs) @@ -4923,10 +4922,11 @@ def prepare_accelerator(args: argparse.Namespace): if args.wandb_api_key is not None: wandb.login(key=args.wandb_api_key) - # torch.compile のオプション。 NO の場合は torch.compile は使わない + # torch.compile If NO, torch.compile is not used." dynamo_backend = "NO" if args.torch_compile: dynamo_backend = args.dynamo_backend + print(f"torch.compile is used with {dynamo_backend} backend") kwargs_handlers = ( InitProcessGroupKwargs(timeout=datetime.timedelta(minutes=args.ddp_timeout)) if args.ddp_timeout else None, diff --git a/nodes.py b/nodes.py index d8314c0..f1d3c29 100644 --- a/nodes.py +++ b/nodes.py @@ -446,7 +446,7 @@ class InitFluxLoRATraining: "t5xxl_max_token_length": 512, "alpha_mask": dataset["alpha_mask"], "network_train_unet_only": True if train_text_encoder == 'disabled' else False, - "fp8_base_unet": False if "fp8" in train_text_encoder else True, + "fp8_base_unet": True if "fp8" in train_text_encoder else False, "disable_mmap_load_safetensors": False, "split_mode": split_mode, } diff --git a/train_db.py b/train_db.py index 8325641..3b8d7b2 100644 --- a/train_db.py +++ b/train_db.py @@ -154,7 +154,7 @@ def train(args): vae.requires_grad_(False) vae.eval() - train_dataset_group.new_cache_latents(vae, accelerator.is_main_process) + train_dataset_group.new_cache_latents(vae, accelerator) vae.to("cpu") clean_memory_on_device(accelerator.device)