allow torch.compiling more stuff

This commit is contained in:
kijai
2024-11-13 23:59:10 +02:00
parent acb19bcd29
commit 723537bb75
4 changed files with 26 additions and 14 deletions
+6 -5
View File
@@ -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))
+11 -5
View File
@@ -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
@@ -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]
@@ -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]