From c5f077451a109b52861191d92ea6228e0dcedec3 Mon Sep 17 00:00:00 2001 From: drozbay <17261091+drozbay@users.noreply.github.com> Date: Wed, 15 Jul 2026 09:21:15 -0600 Subject: [PATCH] Properly offset the Uni3C application for SCAIL. Apply Uni3C across batches. Calculate Uni3C residuals correctly for batched latents. --- comfy/ldm/wan/model.py | 9 +++++++-- comfy_extras/nodes_model_patch.py | 21 ++++++++++++++++----- 2 files changed, 23 insertions(+), 7 deletions(-) diff --git a/comfy/ldm/wan/model.py b/comfy/ldm/wan/model.py index b7ec6f7d2..c042e93c4 100644 --- a/comfy/ldm/wan/model.py +++ b/comfy/ldm/wan/model.py @@ -1698,11 +1698,16 @@ class SCAILWanModel(WanModel): def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs): + x_input = x + + img_offset = 0 if reference_latent is not None: x = torch.cat((reference_latent, x), dim=2) + img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \ + (reference_latent.shape[3] // self.patch_size[1]) * \ + (reference_latent.shape[4] // self.patch_size[2]) # embeddings - x_input = x x = self.patch_embedding(x.float()).to(x.dtype) if ref_mask_latents is not None: # SCAIL-2 additive mask stream (one identity mask frame per reference, then video) x = x + self.patch_embedding_mask(ref_mask_latents.float()).to(x.dtype) @@ -1754,7 +1759,7 @@ class SCAILWanModel(WanModel): if "double_block" in patches: for p in patches["double_block"]: - out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) x = out["img"] # head diff --git a/comfy_extras/nodes_model_patch.py b/comfy_extras/nodes_model_patch.py index 1f83b2e2d..bb8631e76 100644 --- a/comfy_extras/nodes_model_patch.py +++ b/comfy_extras/nodes_model_patch.py @@ -581,9 +581,9 @@ class WanUni3CCnetPatch: comfy.model_management.load_models_gpu(loaded_models) return self.latent_format.process_in(render_latent) - def build_controlnet_input(self, x, dtype): + def build_controlnet_input(self, x, dtype, samples_per_cond): # first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V) - hidden = x[:1, :20].to(dtype) + hidden = x[:samples_per_cond, :20].to(dtype) if hidden.shape[1] < 20: pad_shape = list(hidden.shape) pad_shape[1] = 20 - hidden.shape[1] @@ -594,6 +594,8 @@ class WanUni3CCnetPatch: render = self.encode_render_video(hidden.shape[2:]) render = render.to(device=hidden.device, dtype=dtype) self.prepared_render = render + if render.shape[0] != hidden.shape[0]: + render = render.expand(hidden.shape[0], -1, -1, -1, -1) return torch.cat([hidden, render], dim=1) def __call__(self, kwargs): @@ -610,14 +612,20 @@ class WanUni3CCnetPatch: if sigma > self.sigma_start or sigma < self.sigma_end: active = False if active: - temb = kwargs.get("vec")[:1] + x = kwargs.get("x") + # cond and uncond chunks share latents, so we can reuse residuals + num_conds = len(transformer_options.get("cond_or_uncond", [0])) + samples_per_cond = x.shape[0] + if num_conds > 0 and x.shape[0] % num_conds == 0: + samples_per_cond = x.shape[0] // num_conds + temb = kwargs.get("vec")[:samples_per_cond] if temb.ndim == 3: temb = temb[:, 0] model = self.model_patch.model cnet_dim = model.controlnet_blocks[0].norm1.linear.in_features if temb.shape[-1] != cnet_dim: raise RuntimeError("This Uni3C controlnet expects a Wan model with dim {}, the loaded model has dim {}".format(cnet_dim, temb.shape[-1])) - controlnet_input = self.build_controlnet_input(kwargs.get("x"), img.dtype) + controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond) hidden, freqs = model.process_input(controlnet_input) self.temp_data = (hidden, temb.to(img.dtype), freqs) @@ -625,8 +633,11 @@ class WanUni3CCnetPatch: if self.temp_data is not None and block_index < num_layers: hidden, temb, freqs = self.temp_data hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs) + residual = residual.to(img.dtype) * self.strength + if residual.shape[0] != img.shape[0]: + residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1) img_offset = kwargs.get("img_offset", 0) - img[:, img_offset:img_offset + residual.shape[1]] += residual.to(img.dtype) * self.strength + img[:, img_offset:img_offset + residual.shape[1]] += residual if block_index >= num_layers - 1: self.temp_data = None else: