(
self,
channels,
emb_channels,
dropout,
out_channels=None,
use_conv=False,
use_scale_shift_norm=False,
dims=2,
use_checkpoint=False,
up=False,
down=False,
kernel_size=3,
exchange_temb_dims=False,
skip_t_emb=False,
)
| 223 | """ |
| 224 | |
| 225 | def __init__( |
| 226 | self, |
| 227 | channels, |
| 228 | emb_channels, |
| 229 | dropout, |
| 230 | out_channels=None, |
| 231 | use_conv=False, |
| 232 | use_scale_shift_norm=False, |
| 233 | dims=2, |
| 234 | use_checkpoint=False, |
| 235 | up=False, |
| 236 | down=False, |
| 237 | kernel_size=3, |
| 238 | exchange_temb_dims=False, |
| 239 | skip_t_emb=False, |
| 240 | ): |
| 241 | super().__init__() |
| 242 | self.channels = channels |
| 243 | self.emb_channels = emb_channels |
| 244 | self.dropout = dropout |
| 245 | self.out_channels = out_channels or channels |
| 246 | self.use_conv = use_conv |
| 247 | self.use_checkpoint = use_checkpoint |
| 248 | self.use_scale_shift_norm = use_scale_shift_norm |
| 249 | self.exchange_temb_dims = exchange_temb_dims |
| 250 | |
| 251 | if isinstance(kernel_size, Iterable): |
| 252 | padding = [k // 2 for k in kernel_size] |
| 253 | else: |
| 254 | padding = kernel_size // 2 |
| 255 | |
| 256 | self.in_layers = nn.Sequential( |
| 257 | normalization(channels), |
| 258 | nn.SiLU(), |
| 259 | conv_nd(dims, channels, self.out_channels, kernel_size, padding=padding), |
| 260 | ) |
| 261 | |
| 262 | self.updown = up or down |
| 263 | |
| 264 | if up: |
| 265 | self.h_upd = Upsample(channels, False, dims) |
| 266 | self.x_upd = Upsample(channels, False, dims) |
| 267 | elif down: |
| 268 | self.h_upd = Downsample(channels, False, dims) |
| 269 | self.x_upd = Downsample(channels, False, dims) |
| 270 | else: |
| 271 | self.h_upd = self.x_upd = nn.Identity() |
| 272 | |
| 273 | self.skip_t_emb = skip_t_emb |
| 274 | self.emb_out_channels = 2 * self.out_channels if use_scale_shift_norm else self.out_channels |
| 275 | if self.skip_t_emb: |
| 276 | print(f"Skipping timestep embedding in {self.__class__.__name__}") |
| 277 | assert not self.use_scale_shift_norm |
| 278 | self.emb_layers = None |
| 279 | self.exchange_temb_dims = False |
| 280 | else: |
| 281 | self.emb_layers = nn.Sequential( |
| 282 | nn.SiLU(), |
nothing calls this directly
no test coverage detected