From cf15032ab6c4dbab3e46234753f1795aee175bd4 Mon Sep 17 00:00:00 2001 From: Tai An Date: Mon, 31 Aug 2026 20:21:08 -0700 Subject: [PATCH] fix(pixart): use pe_interpolation instead of the removed lewei_scale in the ControlNet wrappers (#1030) ControlPixArtHalf.forward_c and ControlPixArtMSHalf.forward_raw still use the pre-rename name lewei_scale, both for the attribute and for the get_2d_sincos_pos_embed keyword. Neither exists any more, so both PixArt ControlNet targets fail on the first forward pass. Signed-off-by: Anai-Guo --- py/modules/dit/pixArt/models/pixart_controlnet.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/py/modules/dit/pixArt/models/pixart_controlnet.py b/py/modules/dit/pixArt/models/pixart_controlnet.py index 37fa4c1..75934a3 100644 --- a/py/modules/dit/pixArt/models/pixart_controlnet.py +++ b/py/modules/dit/pixArt/models/pixart_controlnet.py @@ -77,7 +77,7 @@ class ControlPixArtHalf(Module): def forward_c(self, c): self.h, self.w = c.shape[-2]//self.patch_size, c.shape[-1]//self.patch_size - pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), lewei_scale=self.lewei_scale, base_size=self.base_size)).unsqueeze(0).to(c.device).to(self.dtype) + pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation, base_size=self.base_size)).unsqueeze(0).to(c.device).to(self.dtype) return self.x_embedder(c) + pos_embed if c is not None else c # def forward(self, x, t, c, **kwargs): @@ -225,7 +225,7 @@ class ControlPixArtMSHalf(ControlPixArtHalf): c_size, ar = data_info['img_hw'].to(self.dtype), data_info['aspect_ratio'].to(self.dtype) self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size - pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), lewei_scale=self.lewei_scale, base_size=self.base_size)).unsqueeze(0).to(x.device).to(self.dtype) + pos_embed = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation, base_size=self.base_size)).unsqueeze(0).to(x.device).to(self.dtype) x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 t = self.t_embedder(timestep) # (N, D) csize = self.csize_embedder(c_size, bs) # (N, D)