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