From bbda83647da6957e6c0ce52dd86f80b7a6501662 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Thu, 6 Aug 2026 04:12:23 +0300 Subject: [PATCH 1/3] Support int8_convrot VAE (#15334) --- comfy/ldm/minimax/vae.py | 24 +++++++++++++----------- comfy/sd.py | 6 +++++- 2 files changed, 18 insertions(+), 12 deletions(-) diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py index b03bf0c9a23..65d06f3e939 100644 --- a/comfy/ldm/minimax/vae.py +++ b/comfy/ldm/minimax/vae.py @@ -199,11 +199,11 @@ def forward(self, img_ids): class FeedForward(nn.Module): # Gated SiLU FFN. - def __init__(self, dim, mult=4, bias=True): + def __init__(self, dim, mult=4, bias=True, operations=ops): super().__init__() inner_dim = dim * mult - self.w1 = ops.Linear(dim, inner_dim * 2, bias=bias) - self.w2 = ops.Linear(inner_dim, dim, bias=bias) + self.w1 = operations.Linear(dim, inner_dim * 2, bias=bias) + self.w2 = operations.Linear(inner_dim, dim, bias=bias) def forward(self, x): gate, x = self.w1(x).chunk(2, dim=-1) @@ -211,15 +211,15 @@ def forward(self, x): class Attention(nn.Module): - def __init__(self, heads, dim_head, bias=True, eps=1e-5): + def __init__(self, heads, dim_head, bias=True, eps=1e-5, operations=ops): super().__init__() self.dim_head = dim_head self.heads = heads inner_dim = dim_head * heads self.norm_q = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) self.norm_k = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) - self.to_qkv = ops.Linear(inner_dim, inner_dim * 3, bias=bias) - self.to_out = ops.Linear(inner_dim, inner_dim, bias=bias) + self.to_qkv = operations.Linear(inner_dim, inner_dim * 3, bias=bias) + self.to_out = operations.Linear(inner_dim, inner_dim, bias=bias) def forward(self, x, rotary_pos_emb=None): batch_size, seq_len, _ = x.shape @@ -242,14 +242,14 @@ def forward(self, x, rotary_pos_emb=None): class TransformerBlock(nn.Module): - def __init__(self, heads, dim_head, bias=True, eps=1e-5): + def __init__(self, heads, dim_head, bias=True, eps=1e-5, operations=ops): super().__init__() dim = heads * dim_head self.norm1 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) - self.attn = Attention(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + self.attn = Attention(heads=heads, dim_head=dim_head, bias=bias, eps=eps, operations=operations) self.scale1 = nn.Parameter(torch.empty(dim)) self.norm2 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) - self.ff = FeedForward(dim=dim, bias=bias) + self.ff = FeedForward(dim=dim, bias=bias, operations=operations) self.scale2 = nn.Parameter(torch.empty(dim)) def forward(self, x, rotary_pos_emb=None): @@ -259,7 +259,7 @@ def forward(self, x, rotary_pos_emb=None): class ViT3DDecoder(nn.Module): def __init__(self, patch_size=16, patch_size_t=4, in_channels=24, out_channels=3, num_layers=36, heads=32, dim_head=64, rope_theta=100.0, - rope_dim_ratio=0.75, bias=True, eps=1e-5, num_register_tokens=4): + rope_dim_ratio=0.75, bias=True, eps=1e-5, num_register_tokens=4, operations=ops): super().__init__() dim = heads * dim_head self.patch_size = patch_size @@ -274,7 +274,7 @@ def __init__(self, patch_size=16, patch_size_t=4, in_channels=24, out_channels=3 self.register_buffer("mask_token", torch.empty(1, 1, dim)) self.transformer_blocks = nn.ModuleList( - [TransformerBlock(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + [TransformerBlock(heads=heads, dim_head=dim_head, bias=bias, eps=eps, operations=operations) for _ in range(num_layers)] ) @@ -337,6 +337,7 @@ def __init__( tile_size=256, tile_overlap_min=64, tiling=True, + operations=ops, ): super().__init__() self.vae_ratio = int(math.prod(space_down)) @@ -372,6 +373,7 @@ def __init__( patch_size_t=self.vae_ratio_t, in_channels=z_channels, out_channels=out_ch, + operations=operations, ) self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN)) diff --git a/comfy/sd.py b/comfy/sd.py index 8d670106e71..9ccd561bc5c 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -940,7 +940,11 @@ def estimate_memory(shape, dtype, num_layers = 16, kv_cache_multiplier = 2): if not comfy.memory_management.aimdo_enabled: self.disable_offload = True elif "decoder.transformer_blocks.0.scale1" in sd and "encoder.down.5.block.0.conv1.weight" in sd: # MiniMax H3 video VAE - self.first_stage_model = comfy.ldm.minimax.vae.MiniMaxH3VideoVAE() + minimax_ops = comfy.ops.disable_weight_init + minimax_quant = comfy.utils.detect_layer_quantization(sd, "") + if minimax_quant is not None: # int8+convrot quantized decoder + minimax_ops = comfy.ops.mixed_precision_ops(minimax_quant, dtype if dtype is not None else torch.float16) + self.first_stage_model = comfy.ldm.minimax.vae.MiniMaxH3VideoVAE(operations=minimax_ops) self.latent_channels = 24 self.latent_dim = 3 # frames 17k+5 <-> latents 5k+2, 16x spatial From 15989f87ca89bfe2e7c47763252c559e96d97551 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Thu, 6 Aug 2026 04:15:48 +0300 Subject: [PATCH 2/3] Speedup LTX and Wan (#15138) --- comfy/ldm/lightricks/av_model.py | 24 ++++++++++++++++++------ comfy/ldm/lightricks/model.py | 19 ++++++++++++++++--- comfy/ldm/wan/model.py | 10 +++++++++- comfy/ldm/wan/uni3c.py | 4 ++-- 4 files changed, 45 insertions(+), 12 deletions(-) diff --git a/comfy/ldm/lightricks/av_model.py b/comfy/ldm/lightricks/av_model.py index ef993846520..8e360f6a8ba 100644 --- a/comfy/ldm/lightricks/av_model.py +++ b/comfy/ldm/lightricks/av_model.py @@ -16,7 +16,9 @@ from comfy.ldm.lightricks.symmetric_patchifier import AudioPatchifier from comfy.ldm.lightricks.embeddings_connector import Embeddings1DConnector import comfy.ldm.common_dit +import comfy.model_management import comfy.model_prefetch +import comfy.quant_ops class CompressedTimestep: """Store video timestep embeddings in compressed form using per-frame indexing.""" @@ -271,7 +273,10 @@ def forward( if run_vx: # video self-attention vshift_msa, vscale_msa = (self.get_ada_values(self.scale_shift_table, vx.shape[0], v_timestep, slice(0, 2))) - norm_vx = comfy.ldm.common_dit.rms_norm(vx) * (1 + vscale_msa) + vshift_msa + if comfy.model_management.in_training: + norm_vx = comfy.ldm.common_dit.rms_norm(vx) * (1 + vscale_msa) + vshift_msa + else: + norm_vx = comfy.quant_ops.ck.rms_adaln(vx, vscale_msa, vshift_msa) del vshift_msa, vscale_msa attn1_out = self.attn1(norm_vx, pe=v_pe, mask=self_attention_mask, transformer_options=transformer_options) del norm_vx @@ -305,7 +310,6 @@ def forward( # video - audio cross attention. if run_a2v or run_v2a: - vx_norm3 = comfy.ldm.common_dit.rms_norm(vx) ax_norm3 = comfy.ldm.common_dit.rms_norm(ax) # audio to video cross attention @@ -315,7 +319,10 @@ def forward( scale_ca_video_hidden_states_a2v_v, shift_ca_video_hidden_states_a2v_v = self.get_ada_values( self.scale_shift_table_a2v_ca_video[:4, :], vx.shape[0], v_cross_scale_shift_timestep)[:2] - vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_a2v_v) + shift_ca_video_hidden_states_a2v_v + if comfy.model_management.in_training: + vx_scaled = comfy.ldm.common_dit.rms_norm(vx) * (1 + scale_ca_video_hidden_states_a2v_v) + shift_ca_video_hidden_states_a2v_v + else: + vx_scaled = comfy.quant_ops.ck.rms_adaln(vx, scale_ca_video_hidden_states_a2v_v, shift_ca_video_hidden_states_a2v_v) ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_a2v) + shift_ca_audio_hidden_states_a2v del scale_ca_video_hidden_states_a2v_v, shift_ca_video_hidden_states_a2v_v, scale_ca_audio_hidden_states_a2v, shift_ca_audio_hidden_states_a2v @@ -334,7 +341,10 @@ def forward( self.scale_shift_table_a2v_ca_video[:4, :], vx.shape[0], v_cross_scale_shift_timestep)[2:4] ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_v2a) + shift_ca_audio_hidden_states_v2a - vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_v2a) + shift_ca_video_hidden_states_v2a + if comfy.model_management.in_training: + vx_scaled = comfy.ldm.common_dit.rms_norm(vx) * (1 + scale_ca_video_hidden_states_v2a) + shift_ca_video_hidden_states_v2a + else: + vx_scaled = comfy.quant_ops.ck.rms_adaln(vx, scale_ca_video_hidden_states_v2a, shift_ca_video_hidden_states_v2a) del scale_ca_video_hidden_states_v2a, shift_ca_video_hidden_states_v2a, scale_ca_audio_hidden_states_v2a, shift_ca_audio_hidden_states_v2a v2a_out = self.video_to_audio_attn(ax_scaled, context=vx_scaled, pe=a_cross_pe, k_pe=v_cross_pe, transformer_options=transformer_options) @@ -344,12 +354,14 @@ def forward( ax.addcmul_(v2a_out, gate_out_v2a) del gate_out_v2a, v2a_out - del vx_norm3, ax_norm3 # video feedforward if run_vx: vshift_mlp, vscale_mlp = self.get_ada_values(self.scale_shift_table, vx.shape[0], v_timestep, slice(3, 5)) - vx_scaled = comfy.ldm.common_dit.rms_norm(vx) * (1 + vscale_mlp) + vshift_mlp + if comfy.model_management.in_training: + vx_scaled = comfy.ldm.common_dit.rms_norm(vx) * (1 + vscale_mlp) + vshift_mlp + else: + vx_scaled = comfy.quant_ops.ck.rms_adaln(vx, vscale_mlp, vshift_mlp) del vshift_mlp, vscale_mlp ff_out = self.ff(vx_scaled) diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index f9de3a38e24..f80bffba712 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -13,6 +13,7 @@ import comfy.ldm.modules.attention import comfy.ldm.common_dit import comfy.model_management +import comfy.ops import comfy.quant_ops from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords @@ -321,7 +322,11 @@ def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, dtype=None, dev ) def forward(self, x): - return self.net(x) + # net = [GELU_approx(proj), Dropout, Linear]; the fused path skips the + # Dropout, so leave it to the stock path whenever it could be active. + if comfy.model_management.in_training: + return self.net(x) + return comfy.ops.linear_input_act(self.net[2], self.net[0].proj(x), "gelu_tanh") def apply_rotary_emb(input_tensor, freqs_cis): rotation_matrix, split_pe = freqs_cis @@ -535,7 +540,12 @@ def __init__( def forward(self, x, context=None, attention_mask=None, timestep=None, pe=None, transformer_options={}, self_attention_mask=None, prompt_timestep=None): shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (self.scale_shift_table[None, None, :6].to(device=x.device, dtype=x.dtype) + timestep.reshape(x.shape[0], timestep.shape[1], self.scale_shift_table.shape[0], -1)[:, :, :6, :]).unbind(dim=2) - x += self.attn1(comfy.ldm.common_dit.rms_norm(x) * (1 + scale_msa) + shift_msa, pe=pe, mask=self_attention_mask, transformer_options=transformer_options) * gate_msa + if comfy.model_management.in_training: + norm_x = comfy.ldm.common_dit.rms_norm(x) * (1 + scale_msa) + shift_msa + else: + norm_x = comfy.quant_ops.ck.rms_adaln(x, scale_msa, shift_msa) + + x += self.attn1(norm_x, pe=pe, mask=self_attention_mask, transformer_options=transformer_options) * gate_msa if self.cross_attention_adaln: shift_q_mca, scale_q_mca, gate_mca = (self.scale_shift_table[None, None, 6:9].to(device=x.device, dtype=x.dtype) + timestep.reshape(x.shape[0], timestep.shape[1], self.scale_shift_table.shape[0], -1)[:, :, 6:9, :]).unbind(dim=2) @@ -589,7 +599,10 @@ def apply_cross_attention_adaln( prompt_scale_shift_table[None, None].to(device=x.device, dtype=x.dtype) + prompt_timestep.reshape(batch_size, prompt_timestep.shape[1], 2, -1) ).unbind(dim=2) - attn_input = comfy.ldm.common_dit.rms_norm(x) * (1 + q_scale) + q_shift + if comfy.model_management.in_training: + attn_input = comfy.ldm.common_dit.rms_norm(x) * (1 + q_scale) + q_shift + else: + attn_input = comfy.quant_ops.ck.rms_adaln(x, q_scale, q_shift) encoder_hidden_states = context * (1 + scale_kv) + shift_kv return attn(attn_input, context=encoder_hidden_states, mask=attention_mask, transformer_options=transformer_options) * q_gate diff --git a/comfy/ldm/wan/model.py b/comfy/ldm/wan/model.py index c042e93c4c9..dca6efba19d 100644 --- a/comfy/ldm/wan/model.py +++ b/comfy/ldm/wan/model.py @@ -11,6 +11,7 @@ from comfy.ldm.flux.math import apply_rope1, rope import comfy.ldm.common_dit import comfy.model_management +import comfy.ops import comfy.patcher_extension @@ -174,6 +175,13 @@ def repeat_e(e, x): return torch.repeat_interleave(e, repeats + 1, dim=1)[:, :x.size(1)] +class WanFeedForward(nn.Sequential): + """[Linear, GELU(tanh), Linear], with the GELU folded into the down-projection.""" + + def forward(self, x): + return comfy.ops.linear_input_act(self[2], self[0](x), "gelu_tanh") + + class WanAttentionBlock(nn.Module): def __init__(self, @@ -207,7 +215,7 @@ def __init__(self, qk_norm, eps, operation_settings=operation_settings) self.norm2 = operation_settings.get("operations").LayerNorm(dim, eps, elementwise_affine=False, device=operation_settings.get("device"), dtype=operation_settings.get("dtype")) - self.ffn = nn.Sequential( + self.ffn = WanFeedForward( operation_settings.get("operations").Linear(dim, ffn_dim, device=operation_settings.get("device"), dtype=operation_settings.get("dtype")), nn.GELU(approximate='tanh'), operation_settings.get("operations").Linear(ffn_dim, dim, device=operation_settings.get("device"), dtype=operation_settings.get("dtype"))) diff --git a/comfy/ldm/wan/uni3c.py b/comfy/ldm/wan/uni3c.py index 827ad233995..f4bf6820081 100644 --- a/comfy/ldm/wan/uni3c.py +++ b/comfy/ldm/wan/uni3c.py @@ -4,7 +4,7 @@ import torch.nn as nn from comfy.ldm.flux.layers import EmbedND -from .model import WanSelfAttention +from .model import WanFeedForward, WanSelfAttention class Uni3CLayerNormZero(nn.Module): @@ -41,7 +41,7 @@ def __init__( self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings) self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) - self.ffn = nn.Sequential( + self.ffn = WanFeedForward( operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'), operations.Linear(ffn_dim, dim, device=device, dtype=dtype)) From 563b98eefbe643a4cd510ee7f0b43e79880d5a3f Mon Sep 17 00:00:00 2001 From: endman100 Date: Thu, 6 Aug 2026 11:47:13 +0800 Subject: [PATCH 3/3] Fix MiniMax H3 latent noise mask sampling (#15322) --- comfy/model_base.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/comfy/model_base.py b/comfy/model_base.py index 6631c9eb0e5..6f75d53b52d 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -2112,9 +2112,6 @@ def extra_conds(self, **kwargs): out['minimax_payload'] = comfy.conds.CONDConstant(payload) return out - def scale_latent_inpaint(self, sigma, noise, latent_image, **kwargs): - return latent_image - class TripoSplat(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.triposplat.model.LatentSeqMMFlowModel)