(self, img_rgba, img_rgb, padded_img_rgb, use_offset=True)
| 278 | return origin_pixel, y |
| 279 | |
| 280 | def encode(self, img_rgba, img_rgb, padded_img_rgb, use_offset=True): |
| 281 | a_bchw_01 = img_rgba[:, 3:, :, :] |
| 282 | vae_feed = img_rgb.to(device=self.sd_vae.device, dtype=self.sd_vae.dtype) |
| 283 | latent_dist = self.sd_vae.encode(vae_feed).latent_dist |
| 284 | offset_feed = torch.cat([padded_img_rgb, a_bchw_01], dim=1).to(device=self.sd_vae.device, dtype=self.dtype) |
| 285 | offset = self.encoder(offset_feed) * self.alpha |
| 286 | if use_offset: |
| 287 | latent = dist_sample_deterministic(dist=latent_dist, perturbation=offset) |
| 288 | latent = self.sd_vae.config.scaling_factor * (latent - self.sd_vae.config.shift_factor) |
| 289 | else: |
| 290 | latent = latent_dist.sample() |
| 291 | latent = self.sd_vae.config.scaling_factor * (latent - self.sd_vae.config.shift_factor) |
| 292 | return latent |
| 293 | |
| 294 | def forward(self, img_rgba, img_rgb, padded_img_rgb, use_offset=True): |
| 295 | return self.decode(self.encode(img_rgba, img_rgb, padded_img_rgb, use_offset)) |
no test coverage detected