(
self,
in_channels: int,
out_channels: Optional[int] = None,
temb_channels: int = 512,
eps: float = 1e-6,
)
| 552 | """ |
| 553 | |
| 554 | def __init__( |
| 555 | self, |
| 556 | in_channels: int, |
| 557 | out_channels: Optional[int] = None, |
| 558 | temb_channels: int = 512, |
| 559 | eps: float = 1e-6, |
| 560 | ): |
| 561 | super().__init__() |
| 562 | self.in_channels = in_channels |
| 563 | out_channels = in_channels if out_channels is None else out_channels |
| 564 | self.out_channels = out_channels |
| 565 | |
| 566 | kernel_size = (3, 1, 1) |
| 567 | padding = [k // 2 for k in kernel_size] |
| 568 | |
| 569 | self.norm1 = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=eps, affine=True) |
| 570 | self.conv1 = nn.Conv3d( |
| 571 | in_channels, |
| 572 | out_channels, |
| 573 | kernel_size=kernel_size, |
| 574 | stride=1, |
| 575 | padding=padding, |
| 576 | ) |
| 577 | |
| 578 | if temb_channels is not None: |
| 579 | self.time_emb_proj = nn.Linear(temb_channels, out_channels) |
| 580 | else: |
| 581 | self.time_emb_proj = None |
| 582 | |
| 583 | self.norm2 = torch.nn.GroupNorm(num_groups=32, num_channels=out_channels, eps=eps, affine=True) |
| 584 | |
| 585 | self.dropout = torch.nn.Dropout(0.0) |
| 586 | self.conv2 = nn.Conv3d( |
| 587 | out_channels, |
| 588 | out_channels, |
| 589 | kernel_size=kernel_size, |
| 590 | stride=1, |
| 591 | padding=padding, |
| 592 | ) |
| 593 | |
| 594 | self.nonlinearity = get_activation("silu") |
| 595 | |
| 596 | self.use_in_shortcut = self.in_channels != out_channels |
| 597 | |
| 598 | self.conv_shortcut = None |
| 599 | if self.use_in_shortcut: |
| 600 | self.conv_shortcut = nn.Conv3d( |
| 601 | in_channels, |
| 602 | out_channels, |
| 603 | kernel_size=1, |
| 604 | stride=1, |
| 605 | padding=0, |
| 606 | ) |
| 607 | |
| 608 | def forward(self, input_tensor: torch.Tensor, temb: torch.Tensor) -> torch.Tensor: |
| 609 | hidden_states = input_tensor |
nothing calls this directly
no test coverage detected