| 232 | return x |
| 233 | |
| 234 | class TransparentVAE(torch.nn.Module): |
| 235 | def __init__(self, sd_vae, dtype=torch.float16, encoder_file=None, decoder_file=None, alpha=300.0, latent_c=16, *args, **kwargs): |
| 236 | super().__init__(*args, **kwargs) |
| 237 | self.dtype = dtype |
| 238 | |
| 239 | self.sd_vae = sd_vae |
| 240 | self.sd_vae.to(dtype=self.dtype) |
| 241 | self.sd_vae.requires_grad_(False) |
| 242 | |
| 243 | self.encoder = LatentTransparencyOffsetEncoder(latent_c=latent_c) |
| 244 | if encoder_file is not None: |
| 245 | temp = sf.load_file(encoder_file) |
| 246 | # del temp['blocks.16.weight'] |
| 247 | # del temp['blocks.16.bias'] |
| 248 | self.encoder.load_state_dict(temp, strict=True) |
| 249 | del temp |
| 250 | self.encoder.to(dtype=self.dtype) |
| 251 | self.alpha = alpha |
| 252 | |
| 253 | self.decoder = UNet1024(in_channels=3, out_channels=4, latent_c=latent_c) |
| 254 | if decoder_file is not None: |
| 255 | temp = sf.load_file(decoder_file) |
| 256 | # del temp['latent_conv_in.weight'] |
| 257 | # del temp['latent_conv_in.bias'] |
| 258 | self.decoder.load_state_dict(temp, strict=True) |
| 259 | del temp |
| 260 | self.decoder.to(dtype=self.dtype) |
| 261 | self.latent_c = latent_c |
| 262 | |
| 263 | |
| 264 | def sd_decode(self, latent): |
| 265 | return self.sd_vae.decode(latent) |
| 266 | |
| 267 | def decode(self, latent, aug=True): |
| 268 | origin_pixel = self.sd_vae.decode(latent).sample |
| 269 | origin_pixel = (origin_pixel * 0.5 + 0.5) |
| 270 | if not aug: |
| 271 | y = self.decoder(origin_pixel.to(self.dtype), latent.to(self.dtype)) |
| 272 | return origin_pixel, y |
| 273 | list_y = [] |
| 274 | for i in range(int(latent.shape[0])): |
| 275 | y = self.estimate_augmented(origin_pixel[i:i + 1].to(self.dtype), latent[i:i + 1].to(self.dtype)) |
| 276 | list_y.append(y) |
| 277 | y = torch.concat(list_y, dim=0) |
| 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) |
no outgoing calls
no test coverage detected