allow torch.compiling more stuff
This commit is contained in:
+6
-5
@@ -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))
|
||||
|
||||
@@ -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"])
|
||||
@@ -380,6 +383,7 @@ class PyramidFlowVAEEncode:
|
||||
"vae": ("PYRAMIDFLOWVAE",),
|
||||
"image": ("IMAGE",),
|
||||
"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
|
||||
|
||||
@@ -23,6 +23,11 @@ try:
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user