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)