MCPcopy Create free account
hub / github.com/FireRedTeam/LayerDiffuse-Flux / TransparentVAE

Class TransparentVAE

lib_layerdiffuse/vae.py:234–330  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

232 return x
233
234class 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)

Callers 2

demo_i2i.pyFile · 0.90
demo_t2i.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected