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

Method __init__

src/diffusers/models/upsampling.py:92–141  ·  view source on GitHub ↗
(
        self,
        channels: int,
        use_conv: bool = False,
        use_conv_transpose: bool = False,
        out_channels: Optional[int] = None,
        name: str = "conv",
        kernel_size: Optional[int] = None,
        padding=1,
        norm_type=None,
        eps=None,
        elementwise_affine=None,
        bias=True,
        interpolate=True,
    )

Source from the content-addressed store, hash-verified

90 """
91
92 def __init__(
93 self,
94 channels: int,
95 use_conv: bool = False,
96 use_conv_transpose: bool = False,
97 out_channels: Optional[int] = None,
98 name: str = "conv",
99 kernel_size: Optional[int] = None,
100 padding=1,
101 norm_type=None,
102 eps=None,
103 elementwise_affine=None,
104 bias=True,
105 interpolate=True,
106 ):
107 super().__init__()
108 self.channels = channels
109 self.out_channels = out_channels or channels
110 self.use_conv = use_conv
111 self.use_conv_transpose = use_conv_transpose
112 self.name = name
113 self.interpolate = interpolate
114 conv_cls = nn.Conv2d if USE_PEFT_BACKEND else LoRACompatibleConv
115
116 if norm_type == "ln_norm":
117 self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
118 elif norm_type == "rms_norm":
119 self.norm = RMSNorm(channels, eps, elementwise_affine)
120 elif norm_type is None:
121 self.norm = None
122 else:
123 raise ValueError(f"unknown norm_type: {norm_type}")
124
125 conv = None
126 if use_conv_transpose:
127 if kernel_size is None:
128 kernel_size = 4
129 conv = nn.ConvTranspose2d(
130 channels, self.out_channels, kernel_size=kernel_size, stride=2, padding=padding, bias=bias
131 )
132 elif use_conv:
133 if kernel_size is None:
134 kernel_size = 3
135 conv = conv_cls(self.channels, self.out_channels, kernel_size=kernel_size, padding=padding, bias=bias)
136
137 # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed
138 if name == "conv":
139 self.conv = conv
140 else:
141 self.Conv2d_0 = conv
142
143 def forward(
144 self,

Callers

nothing calls this directly

Calls 2

RMSNormClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected