(self, sd_vae, dtype=torch.float16, encoder_file=None, decoder_file=None, alpha=300.0, latent_c=16, *args, **kwargs)
| 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): |
nothing calls this directly
no test coverage detected