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,)的张量,但问题依旧。
问题原因
- 输入张量重复拼接逻辑错误:文本条件分支中,每次循环都执行
input = torch.cat([input] * bs),当批量>1时,第一次循环后input的批量会变为原批量×bs,第二次循环又会再乘一次bs,导致input批量指数级增长,和denoise_kwargs里的encoder_hidden_states批量不匹配,最终UNet输出张量维度混乱。 - 引导判断逻辑硬编码失效:代码用
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
相关产品推荐
相关产品推荐

