small fix

This commit is contained in:
kijai
2024-11-07 00:17:03 +02:00
parent 7932a63b86
commit 8bfe1f668b
6 changed files with 8 additions and 8 deletions
+1 -1
View File
@@ -175,7 +175,7 @@ def train(args):
vae.requires_grad_(False) vae.requires_grad_(False)
vae.eval() vae.eval()
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process) train_dataset_group.new_cache_latents(vae, accelerator)
vae.to("cpu") vae.to("cpu")
clean_memory_on_device(accelerator.device) clean_memory_on_device(accelerator.device)
+1 -1
View File
@@ -181,7 +181,7 @@ class FluxTrainer:
ae.requires_grad_(False) ae.requires_grad_(False)
ae.eval() 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 ae.to("cpu") # if no sampling, vae can be deleted
clean_memory_on_device(accelerator.device) clean_memory_on_device(accelerator.device)
+2 -2
View File
@@ -438,11 +438,11 @@ configs = {
# region math # region math
from sageattention import sageattn
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, attn_mask: Optional[Tensor] = None) -> Tensor: def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, attn_mask: Optional[Tensor] = None) -> Tensor:
q, k = apply_rope(q, k, pe) 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)") x = rearrange(x, "B H L D -> B L (H D)")
return x return x
+2 -2
View File
@@ -1507,7 +1507,6 @@ class BaseDataset(torch.utils.data.Dataset):
text_encoder_outputs = self.text_encoder_output_caching_strategy.load_outputs_npz( text_encoder_outputs = self.text_encoder_output_caching_strategy.load_outputs_npz(
image_info.text_encoder_outputs_npz image_info.text_encoder_outputs_npz
) )
text_encoder_outputs = [torch.FloatTensor(x) for x in text_encoder_outputs]
else: else:
tokenization_required = True tokenization_required = True
text_encoder_outputs_list.append(text_encoder_outputs) 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: if args.wandb_api_key is not None:
wandb.login(key=args.wandb_api_key) wandb.login(key=args.wandb_api_key)
# torch.compile のオプション。 NO の場合は torch.compile は使わない # torch.compile If NO, torch.compile is not used."
dynamo_backend = "NO" dynamo_backend = "NO"
if args.torch_compile: if args.torch_compile:
dynamo_backend = args.dynamo_backend dynamo_backend = args.dynamo_backend
print(f"torch.compile is used with {dynamo_backend} backend")
kwargs_handlers = ( kwargs_handlers = (
InitProcessGroupKwargs(timeout=datetime.timedelta(minutes=args.ddp_timeout)) if args.ddp_timeout else None, InitProcessGroupKwargs(timeout=datetime.timedelta(minutes=args.ddp_timeout)) if args.ddp_timeout else None,
+1 -1
View File
@@ -446,7 +446,7 @@ class InitFluxLoRATraining:
"t5xxl_max_token_length": 512, "t5xxl_max_token_length": 512,
"alpha_mask": dataset["alpha_mask"], "alpha_mask": dataset["alpha_mask"],
"network_train_unet_only": True if train_text_encoder == 'disabled' else False, "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, "disable_mmap_load_safetensors": False,
"split_mode": split_mode, "split_mode": split_mode,
} }
+1 -1
View File
@@ -154,7 +154,7 @@ def train(args):
vae.requires_grad_(False) vae.requires_grad_(False)
vae.eval() vae.eval()
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process) train_dataset_group.new_cache_latents(vae, accelerator)
vae.to("cpu") vae.to("cpu")
clean_memory_on_device(accelerator.device) clean_memory_on_device(accelerator.device)