From 723537bb75316de1dea96f17405dcc599f624c41 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 13 Nov 2024 23:59:10 +0200 Subject: [PATCH] allow torch.compiling more stuff --- fp8_optimization.py | 11 ++++++----- nodes.py | 16 +++++++++++----- pyramid_dit/flux_modules/modeling_flux_block.py | 12 +++++++++--- .../pyramid_dit_for_video_gen_pipeline.py | 1 - 4 files changed, 26 insertions(+), 14 deletions(-) diff --git a/fp8_optimization.py b/fp8_optimization.py index b01ac91..cd416ce 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -36,10 +36,11 @@ def fp8_linear_forward(cls, original_dtype, input): else: return cls.original_forward(input) -def convert_fp8_linear(module, original_dtype): +def convert_fp8_linear(module, original_dtype, params_to_keep): setattr(module, "fp8_matmul_enabled", True) for name, module in module.named_modules(): - if isinstance(module, nn.Linear): - original_forward = module.forward - setattr(module, "original_forward", original_forward) - setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input)) + if not any(keyword in name for keyword in params_to_keep): + if isinstance(module, nn.Linear): + original_forward = module.forward + setattr(module, "original_forward", original_forward) + setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input)) diff --git a/nodes.py b/nodes.py index 3ae2a44..9a9a99b 100644 --- a/nodes.py +++ b/nodes.py @@ -43,6 +43,7 @@ class PyramidFlowTorchCompileSettings: "double_blocks": ("BOOLEAN", {"default": True, "tooltip": "Compile transformer blocks"}), "embedders": ("BOOLEAN", {"default": True, "tooltip": "Compile embedders"}), "compile_rest": ("BOOLEAN", {"default": True, "tooltip": "Compile the rest of the model (proj and norm out)"}), + "dynamo_cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.cache_size_limit"}), }, } RETURN_TYPES = ("PYRAMIDFLOW_COMPILEARGS",) @@ -51,7 +52,7 @@ class PyramidFlowTorchCompileSettings: CATEGORY = "MochiWrapper" DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended" - def loadmodel(self, backend, fullgraph, mode, compile_whole_model, single_blocks, double_blocks, embedders, compile_rest): + def loadmodel(self, backend, fullgraph, mode, compile_whole_model, single_blocks, double_blocks, embedders, compile_rest, dynamo_cache_size_limit): compile_args = { "backend": backend, @@ -62,6 +63,7 @@ class PyramidFlowTorchCompileSettings: "double_blocks": double_blocks, "embedders": embedders, "compile_rest": compile_rest, + "dynamo_cache_size_limit": dynamo_cache_size_limit, } return (compile_args, ) @@ -175,7 +177,7 @@ class PyramidFlowModelLoader: with open(config_path) as f: config = json.load(f) transformer = PyramidDiffusionMMDiT.from_config(config) - params_to_keep = {"pos_embedding"} + params_to_keep = {"pos_embedding", "norm_k", "norm_q", "norm_v", "norm_added_k", "norm_added_q", "bias"} if is_accelerate_available: logging.info("Using accelerate to load and assign model weights to device...") for name, param in transformer.named_parameters(): @@ -192,12 +194,13 @@ class PyramidFlowModelLoader: if precision == "fp8_e4m3fn_fast": from .fp8_optimization import convert_fp8_linear - convert_fp8_linear(transformer, torch.bfloat16) + convert_fp8_linear(transformer, torch.bfloat16, params_to_keep=params_to_keep) transformer.to(device) #torch.compile if compile_args is not None: torch._dynamo.config.force_parameter_static_shapes = False + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] dynamic = True # because of the stages the compiliation should be dynamic if compile_args["compile_whole_model"]: transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"]) @@ -379,7 +382,8 @@ class PyramidFlowVAEEncode: "required": { "vae": ("PYRAMIDFLOWVAE",), "image": ("IMAGE",), - "enable_tiling": ("BOOLEAN", {"default": False}), + "enable_tiling": ("BOOLEAN", {"default": False}), + "overlap_factor": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}), }, } @@ -388,7 +392,7 @@ class PyramidFlowVAEEncode: FUNCTION = "sample" CATEGORY = "PyramidFlowWrapper" - def sample(self, vae, image, enable_tiling): + def sample(self, vae, image, enable_tiling, overlap_factor): B, H, W, C = image.shape mm.soft_empty_cache() @@ -400,6 +404,8 @@ class PyramidFlowVAEEncode: device = mm.get_torch_device() offload_device = mm.unet_offload_device() + vae.encode_tile_overlap_factor = overlap_factor + # For the image latent vae_shift_factor = 0.1490 vae_scale_factor = 1 / 1.8415 diff --git a/pyramid_dit/flux_modules/modeling_flux_block.py b/pyramid_dit/flux_modules/modeling_flux_block.py index 1989698..a55d5e7 100644 --- a/pyramid_dit/flux_modules/modeling_flux_block.py +++ b/pyramid_dit/flux_modules/modeling_flux_block.py @@ -22,7 +22,12 @@ try: from flash_attn.flash_attn_interface import flash_attn_varlen_func except: flash_attn_varlen_func = None - + +@torch.compiler.disable() +def compute_attention(query, key, value, attn_mask, dropout_p=0.0, is_causal=False): + return F.scaled_dot_product_attention( + query, key, value, dropout_p=dropout_p, is_causal=is_causal, attn_mask=attn_mask, + ) def apply_rope(xq, xk, freqs_cis): xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2) @@ -204,7 +209,7 @@ class VarlenSelfAttentionWithT5Mask: value = value.transpose(1, 2) # with torch.backends.cuda.sdp_kernel(enable_math=False, enable_flash=False, enable_mem_efficient=True): - stage_hidden_states = F.scaled_dot_product_attention( + stage_hidden_states = compute_attention( query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p], ) stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim] @@ -219,6 +224,7 @@ class VarlenSelfAttentionWithT5Mask: return output_hidden, output_encoder_hidden + class VarlenFlashSelfAttnSingle: def __init__(self): @@ -312,7 +318,7 @@ class VarlenSelfAttnSingle: key = key.transpose(1, 2).contiguous() value = value.transpose(1, 2).contiguous() - stage_hidden_states = F.scaled_dot_product_attention( + stage_hidden_states = compute_attention( query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p], ) stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim] diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index d9862b7..6b559d5 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -150,7 +150,6 @@ class PyramidDiTForVideoGeneration: latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) return latents - def sample_block_noise(self, bs, ch, temp, height, width): block_number = bs * ch * temp * (height // 2) * (width // 2) noise = torch.stack([self.dist.sample() for _ in range(block_number)]) # [block number, 4]