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.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)
+1 -1
View File
@@ -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)
+2 -2
View File
@@ -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
+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(
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,
+1 -1
View File
@@ -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,
}
+1 -1
View File
@@ -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)