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

Class TransparentVAEEncoder

lib_layerdiffuse/vae.py:414–447  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

412
413
414class TransparentVAEEncoder(torch.nn.Module):
415 def __init__(self, filename, dtype=torch.float16, alpha=300.0, *args, **kwargs):
416 super().__init__(*args, **kwargs)
417 sd = sf.load_file(filename)
418 self.dtype = dtype
419
420 model = LatentTransparencyOffsetEncoder()
421 model.load_state_dict(sd, strict=True)
422 model.to(dtype=self.dtype)
423 model.eval()
424
425 self.model = model
426
427 # similar to LoRA's alpha to avoid initial zero-initialized outputs being too small
428 self.alpha = alpha
429 return
430
431 @torch.no_grad()
432 def forward(self, sd_vae, list_of_np_rgba_hwc_uint8, use_offset=True):
433 list_of_np_rgb_padded = [pad_rgb(x) for x in list_of_np_rgba_hwc_uint8]
434 rgb_padded_bchw_01 = torch.from_numpy(np.stack(list_of_np_rgb_padded, axis=0)).float().movedim(-1, 1)
435 rgba_bchw_01 = torch.from_numpy(np.stack(list_of_np_rgba_hwc_uint8, axis=0)).float().movedim(-1, 1) / 255.0
436 rgb_bchw_01 = rgba_bchw_01[:, :3, :, :]
437 a_bchw_01 = rgba_bchw_01[:, 3:, :, :]
438 vae_feed = (rgb_bchw_01 * 2.0 - 1.0) * a_bchw_01
439 vae_feed = vae_feed.to(device=sd_vae.device, dtype=sd_vae.dtype)
440 latent_dist = sd_vae.encode(vae_feed).latent_dist
441 offset_feed = torch.cat([a_bchw_01, rgb_padded_bchw_01], dim=1).to(device=sd_vae.device, dtype=self.dtype)
442 offset = self.model(offset_feed) * self.alpha
443 if use_offset:
444 latent = dist_sample_deterministic(dist=latent_dist, perturbation=offset)
445 else:
446 latent = latent_dist.sample()
447 return latent

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected