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

Method __init__

src/diffusers/models/vq_model.py:75–127  ·  view source on GitHub ↗
(
        self,
        in_channels: int = 3,
        out_channels: int = 3,
        down_block_types: Tuple[str, ...] = ("DownEncoderBlock2D",),
        up_block_types: Tuple[str, ...] = ("UpDecoderBlock2D",),
        block_out_channels: Tuple[int, ...] = (64,),
        layers_per_block: int = 1,
        act_fn: str = "silu",
        latent_channels: int = 3,
        sample_size: int = 32,
        num_vq_embeddings: int = 256,
        norm_num_groups: int = 32,
        vq_embed_dim: Optional[int] = None,
        scaling_factor: float = 0.18215,
        norm_type: str = "group",  # group, spatial
        mid_block_add_attention=True,
        lookup_from_codebook=False,
        force_upcast=False,
    )

Source from the content-addressed store, hash-verified

73
74 @register_to_config
75 def __init__(
76 self,
77 in_channels: int = 3,
78 out_channels: int = 3,
79 down_block_types: Tuple[str, ...] = ("DownEncoderBlock2D",),
80 up_block_types: Tuple[str, ...] = ("UpDecoderBlock2D",),
81 block_out_channels: Tuple[int, ...] = (64,),
82 layers_per_block: int = 1,
83 act_fn: str = "silu",
84 latent_channels: int = 3,
85 sample_size: int = 32,
86 num_vq_embeddings: int = 256,
87 norm_num_groups: int = 32,
88 vq_embed_dim: Optional[int] = None,
89 scaling_factor: float = 0.18215,
90 norm_type: str = "group", # group, spatial
91 mid_block_add_attention=True,
92 lookup_from_codebook=False,
93 force_upcast=False,
94 ):
95 super().__init__()
96
97 # pass init params to Encoder
98 self.encoder = Encoder(
99 in_channels=in_channels,
100 out_channels=latent_channels,
101 down_block_types=down_block_types,
102 block_out_channels=block_out_channels,
103 layers_per_block=layers_per_block,
104 act_fn=act_fn,
105 norm_num_groups=norm_num_groups,
106 double_z=False,
107 mid_block_add_attention=mid_block_add_attention,
108 )
109
110 vq_embed_dim = vq_embed_dim if vq_embed_dim is not None else latent_channels
111
112 self.quant_conv = nn.Conv2d(latent_channels, vq_embed_dim, 1)
113 self.quantize = VectorQuantizer(num_vq_embeddings, vq_embed_dim, beta=0.25, remap=None, sane_index_shape=False)
114 self.post_quant_conv = nn.Conv2d(vq_embed_dim, latent_channels, 1)
115
116 # pass init params to Decoder
117 self.decoder = Decoder(
118 in_channels=latent_channels,
119 out_channels=out_channels,
120 up_block_types=up_block_types,
121 block_out_channels=block_out_channels,
122 layers_per_block=layers_per_block,
123 act_fn=act_fn,
124 norm_num_groups=norm_num_groups,
125 norm_type=norm_type,
126 mid_block_add_attention=mid_block_add_attention,
127 )
128
129 @apply_forward_hook
130 def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> VQEncoderOutput:

Callers

nothing calls this directly

Calls 3

EncoderClass · 0.85
VectorQuantizerClass · 0.85
DecoderClass · 0.85

Tested by

no test coverage detected