(self, z, temb=None)
| 723 | padding=1) |
| 724 | |
| 725 | def forward(self, z, temb=None): |
| 726 | #assert z.shape[1:] == self.z_shape[1:] |
| 727 | outputs = {} |
| 728 | self.last_z_shape = z.shape |
| 729 | |
| 730 | # timestep embedding |
| 731 | #temb = None |
| 732 | |
| 733 | # z to block_in |
| 734 | h = self.conv_in(z) |
| 735 | if 'conv_in' in self.output_key: |
| 736 | outputs['conv_in'] = h.clone() |
| 737 | |
| 738 | # middle |
| 739 | h = self.mid.block_1(h, temb) |
| 740 | # h = self.mid.attn_1(h) |
| 741 | h = self.mid.block_2(h, temb) |
| 742 | |
| 743 | if 'mid' in self.output_key: |
| 744 | outputs['mid'] = h.clone() |
| 745 | |
| 746 | # upsampling |
| 747 | for i_level in reversed(range(self.num_resolutions)): |
| 748 | for i_block in range(self.num_res_blocks+1): |
| 749 | h = self.up[i_level].block[i_block](h, temb) |
| 750 | if len(self.up[i_level].attn) > 0: |
| 751 | h = self.up[i_level].attn[i_block](h) |
| 752 | if i_level != 0: |
| 753 | h = self.up[i_level].upsample(h) |
| 754 | |
| 755 | block_idx = self.num_resolutions - i_level - 1 |
| 756 | if f'up_block{block_idx}' in self.output_key: |
| 757 | outputs[f'up_block{block_idx}'] = h.clone() |
| 758 | |
| 759 | # end |
| 760 | if self.give_pre_end: |
| 761 | return h |
| 762 | |
| 763 | h = self.norm_out(h) |
| 764 | if 'norm_out' in self.output_key: |
| 765 | outputs['norm_out'] = h.clone() |
| 766 | |
| 767 | h = nonlinearity(h) |
| 768 | if 'nonlinear' in self.output_key: |
| 769 | outputs['nonlinear'] = h.clone() |
| 770 | |
| 771 | h = self.conv_out(h) |
| 772 | # outputs['final'] = h |
| 773 | return h |
| 774 | |
| 775 | |
| 776 | class UNetBlockDecoder(nn.Module): |
nothing calls this directly
no test coverage detected