small fix
This commit is contained in:
+1
-1
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user