(self, x, z)
| 650 | |
| 651 | |
| 652 | def forward(self, x, z): |
| 653 | #assert x.shape[2] == x.shape[3] == self.resolution |
| 654 | |
| 655 | if self.use_timestep: |
| 656 | # timestep embedding |
| 657 | assert t is not None |
| 658 | temb = get_timestep_embedding(t, self.ch) |
| 659 | temb = self.temb.dense[0](temb) |
| 660 | temb = nonlinearity(temb) |
| 661 | temb = self.temb.dense[1](temb) |
| 662 | else: |
| 663 | temb = None |
| 664 | |
| 665 | # downsampling |
| 666 | hs = [self.conv_in(x)] |
| 667 | for i_level in range(self.num_resolutions): |
| 668 | for i_block in range(self.num_res_blocks): |
| 669 | h = self.down[i_level].block[i_block](hs[-1], temb) |
| 670 | if len(self.down[i_level].attn) > 0: |
| 671 | h = self.down[i_level].attn[i_block](h) |
| 672 | hs.append(h) |
| 673 | if i_level != self.num_resolutions-1: |
| 674 | hs.append(self.down[i_level].downsample(hs[-1])) |
| 675 | |
| 676 | # middle |
| 677 | h = hs[-1] |
| 678 | z = self.z_in(z) |
| 679 | h = torch.cat((h,z),dim=1) |
| 680 | h = self.mid.block_1(h, temb) |
| 681 | h = self.mid.attn_1(h) |
| 682 | h = self.mid.block_2(h, temb) |
| 683 | |
| 684 | # upsampling |
| 685 | for i_level in reversed(range(self.num_resolutions)): |
| 686 | for i_block in range(self.num_res_blocks+1): |
| 687 | h = self.up[i_level].block[i_block]( |
| 688 | torch.cat([h, hs.pop()], dim=1), temb) |
| 689 | if len(self.up[i_level].attn) > 0: |
| 690 | h = self.up[i_level].attn[i_block](h) |
| 691 | if i_level != 0: |
| 692 | h = self.up[i_level].upsample(h) |
| 693 | |
| 694 | # end |
| 695 | h = self.norm_out(h) |
| 696 | h = nonlinearity(h) |
| 697 | h = self.conv_out(h) |
| 698 | return h |
| 699 | |
| 700 | |
| 701 | class SimpleDecoder(nn.Module): |
nothing calls this directly
no test coverage detected