You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Stable Diffusion DDIM反演时批量>1的UNet张量维度不匹配问题

DDIM反演Stable Diffusion批量大于1时的张量维度不匹配问题

问题描述

在使用DDIM反演代码运行Stable Diffusion模型时,当批量大小设为大于1(比如32)时,触发以下错误:

RuntimeError: The size of tensor a (131072) must match the size of tensor b (4096) at non-singleton dimension 1.

131072是32×4096的结果,明确属于张量维度不匹配,错误出在这行代码:

noisy_residual = self.unet(input, t.to(input.device), **denoise_kwargs).sample

相关反演代码片段:

## Inversion
def invert_process(self, guidance_scale, input, denoise_kwargs):

    pred_images = []
    pred_latents = []
    
    decode_kwargs = {'vae': self.vae}

    # Reversed timesteps
    timesteps = reversed(self.scheduler.timesteps)
    num_inference_steps = len(self.scheduler.timesteps)

    with torch.no_grad():
        for i in tqdm(range(0, num_inference_steps)):

            t = timesteps[i]
            self.cur_t = t.item()
            
            # For text condition on stable diffusion
            if 'encoder_hidden_states' in denoise_kwargs.keys():
                bs = denoise_kwargs['encoder_hidden_states'].shape[0]
                input = torch.cat([input] * bs)

            # Predict the noise residual
            noisy_residual = self.unet(input, t.to(input.device), **denoise_kwargs).sample
            noise_pred = noisy_residual

            # For text condition on stable diffusion
            if noisy_residual.shape[0] == 2:
                # perform guidance
                noise_pred_text, noise_pred_uncond = noisy_residual.chunk(2)
                noisy_residual = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
                input, _ = input.chunk(2)

            current_t = max(0, self.cur_t - (1000//num_inference_steps)) #t
            next_t = t # min(999, t.item() + (1000//num_inference_steps)) # t+1
            alpha_t = self.scheduler.alphas_cumprod[current_t].to(self.device)
            alpha_t_next = self.scheduler.alphas_cumprod[next_t].to(self.device)

            latents = input

            # Inverted update step (re-arranging the update step to get x(t) (new latents) as a function of x(t-1) (current latents)
            # Add noise to latents

            latents = (latents - (1-alpha_t).sqrt()*noise_pred)*(alpha_t_next.sqrt()/alpha_t.sqrt()) + (1-alpha_t_next).sqrt()*noise_pred
            
            input = latents
            
            pred_latents.append(latents)
            pred_images.append(decode_latent(latents, **decode_kwargs))
            
    return pred_images, pred_latents

已确认批量大小为1时模型运行正常,尝试过将t改为形状为(batch size,)的张量,但问题依旧。


问题原因

  1. 输入张量重复拼接逻辑错误:文本条件分支中,每次循环都执行input = torch.cat([input] * bs),当批量>1时,第一次循环后input的批量会变为原批量×bs,第二次循环又会再乘一次bs,导致input批量指数级增长,和denoise_kwargs里的encoder_hidden_states批量不匹配,最终UNet输出张量维度混乱。
  2. 引导判断逻辑硬编码失效:代码用if noisy_residual.shape[0] == 2:判断是否需要引导,但批量>1时,带引导的UNet输出批量应为2×原批量,而非固定值2,这个判断完全不适用,导致后续chunk(2)拆分出的张量维度错误,引发后续计算的维度不匹配。

解决方法

1. 修正输入拼接逻辑,仅处理一次

将输入拼接逻辑移到循环外,避免重复拼接:

def invert_process(self, guidance_scale, input, denoise_kwargs):

    pred_images = []
    pred_latents = []
    
    decode_kwargs = {'vae': self.vae}
    timesteps = reversed(self.scheduler.timesteps)
    num_inference_steps = len(self.scheduler.timesteps)

    # 提前处理文本条件的输入拼接,仅执行一次
    use_guidance = 'encoder_hidden_states' in denoise_kwargs.keys()
    original_bs = input.shape[0]
    if use_guidance:
        # 拼接无条件和有条件的latent,对应文本条件的拼接
        input = torch.cat([input] * 2)

    with torch.no_grad():
        for i in tqdm(range(0, num_inference_steps)):

            t = timesteps[i]
            self.cur_t = t.item()
            # 适配时间步的批量维度
            t = t.expand(input.shape[0])
            
            # 预测噪声残差
            noisy_residual = self.unet(input, t.to(input.device), **denoise_kwargs).sample

            # 处理引导逻辑,适配批量>1的情况
            if use_guidance:
                # 按批量维度拆分为无条件和有条件的预测
                noise_pred_uncond, noise_pred_text = noisy_residual.chunk(2)
                noisy_residual = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
                # 用回原始批量的latent进行后续计算
                current_input = input.chunk(2)[0]
                noise_pred = noise_pred_uncond
            else:
                current_input = input
                noise_pred = noisy_residual

            current_t = max(0, self.cur_t - (1000//num_inference_steps))
            next_t = t[0].item()  # 取单个时间步值即可
            alpha_t = self.scheduler.alphas_cumprod[current_t].to(self.device)
            alpha_t_next = self.scheduler.alphas_cumprod[next_t].to(self.device)

            latents = current_input
            # 反演更新步骤
            latents = (latents - (1-alpha_t).sqrt()*noise_pred)*(alpha_t_next.sqrt()/alpha_t.sqrt()) + (1-alpha_t_next).sqrt()*noise_pred
            
            # 准备下一轮的输入:如果用引导,重新拼接latent
            if use_guidance:
                input = torch.cat([latents] * 2)
            else:
                input = latents
            
            pred_latents.append(latents)
            pred_images.append(decode_latent(latents, **decode_kwargs))
            
    return pred_images, pred_latents

2. 验证维度一致性(可选调试步骤)

在UNet调用前添加维度检查,确保输入和条件张量批量匹配:

if use_guidance:
    print(f"Input batch: {input.shape[0]}, Encoder states batch: {denoise_kwargs['encoder_hidden_states'].shape[0]}")

若使用引导,两者批量应均为2×原批量。


内容的提问来源于stack exchange,提问作者young

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 10:38:09