MCPcopy Create free account
hub / github.com/CompVis/diff2flow / __init__

Method __init__

diff2flow/kl_autoencoder.py:464–492  ·  view source on GitHub ↗
(
            self,
            ckpt_path: str = None,
            ddconfig=DEFAULT_DDCONFIG,
            embed_dim: int = 4,
            scale: float = 0.18215,     # SD: 0.18215
            shift: float = 0.0
        )

Source from the content-addressed store, hash-verified

462
463class AutoencoderKL(nn.Module):
464 def __init__(
465 self,
466 ckpt_path: str = None,
467 ddconfig=DEFAULT_DDCONFIG,
468 embed_dim: int = 4,
469 scale: float = 0.18215, # SD: 0.18215
470 shift: float = 0.0
471 ):
472 super().__init__()
473 self.encoder = Encoder(**ddconfig)
474 self.decoder = Decoder(**ddconfig)
475 assert ddconfig["double_z"]
476 self.quant_conv = nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
477 self.post_quant_conv = nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
478 self.embed_dim = embed_dim
479
480 self.scale = scale
481 self.shift = shift
482
483 if exists(ckpt_path):
484 assert os.path.exists(ckpt_path), f'[AutoencoderKL] Checkpoint {ckpt_path} not found!'
485 print(f'[AutoencoderKL] Loading checkpoint from {ckpt_path}')
486 if torch.cuda.is_available():
487 self.load_state_dict(torch.load(ckpt_path, weights_only=True))
488 else:
489 self.load_state_dict(torch.load(ckpt_path, weights_only=True, map_location=torch.device('cpu')))
490 else:
491 import warnings
492 warnings.warn(f'[AutoencoderKL] No checkpoint provided. Random initialization.')
493
494 @torch.no_grad()
495 def encode(self, x: torch.Tensor, return_posterior=False):

Callers

nothing calls this directly

Calls 4

EncoderClass · 0.70
DecoderClass · 0.70
existsFunction · 0.70
__init__Method · 0.45

Tested by

no test coverage detected