MCPcopy Create free account
hub / github.com/TencentARC/BrushNet / __init__

Method __init__

src/diffusers/models/autoencoders/vae.py:656–691  ·  view source on GitHub ↗
(
        self,
        n_e: int,
        vq_embed_dim: int,
        beta: float,
        remap=None,
        unknown_index: str = "random",
        sane_index_shape: bool = False,
        legacy: bool = True,
    )

Source from the content-addressed store, hash-verified

654 # backwards compatibility we use the buggy version by default, but you can
655 # specify legacy=False to fix it.
656 def __init__(
657 self,
658 n_e: int,
659 vq_embed_dim: int,
660 beta: float,
661 remap=None,
662 unknown_index: str = "random",
663 sane_index_shape: bool = False,
664 legacy: bool = True,
665 ):
666 super().__init__()
667 self.n_e = n_e
668 self.vq_embed_dim = vq_embed_dim
669 self.beta = beta
670 self.legacy = legacy
671
672 self.embedding = nn.Embedding(self.n_e, self.vq_embed_dim)
673 self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e)
674
675 self.remap = remap
676 if self.remap is not None:
677 self.register_buffer("used", torch.tensor(np.load(self.remap)))
678 self.used: torch.Tensor
679 self.re_embed = self.used.shape[0]
680 self.unknown_index = unknown_index # "random" or "extra" or integer
681 if self.unknown_index == "extra":
682 self.unknown_index = self.re_embed
683 self.re_embed = self.re_embed + 1
684 print(
685 f"Remapping {self.n_e} indices to {self.re_embed} indices. "
686 f"Using {self.unknown_index} for unknown indices."
687 )
688 else:
689 self.re_embed = n_e
690
691 self.sane_index_shape = sane_index_shape
692
693 def remap_to_used(self, inds: torch.LongTensor) -> torch.LongTensor:
694 ishape = inds.shape

Callers

nothing calls this directly

Calls 2

__init__Method · 0.45
loadMethod · 0.45

Tested by

no test coverage detected