(self, z, clear_fake_cp_cache=True)
| 961 | print("Decoder3D initialized.") |
| 962 | |
| 963 | def forward(self, z, clear_fake_cp_cache=True): |
| 964 | self.last_z_shape = z.shape |
| 965 | |
| 966 | # timestep embedding |
| 967 | temb = None |
| 968 | |
| 969 | t = z.shape[2] |
| 970 | # z to block_in |
| 971 | |
| 972 | zq = z |
| 973 | h = self.conv_in(z, clear_cache=clear_fake_cp_cache) |
| 974 | |
| 975 | # middle |
| 976 | h = self.mid.block_1(h, temb, zq, clear_fake_cp_cache=clear_fake_cp_cache) |
| 977 | h = self.mid.block_2(h, temb, zq, clear_fake_cp_cache=clear_fake_cp_cache) |
| 978 | |
| 979 | # upsampling |
| 980 | for i_level in reversed(range(self.num_resolutions)): |
| 981 | for i_block in range(self.num_res_blocks + 1): |
| 982 | h = self.up[i_level].block[i_block](h, temb, zq, clear_fake_cp_cache=clear_fake_cp_cache) |
| 983 | if len(self.up[i_level].attn) > 0: |
| 984 | h = self.up[i_level].attn[i_block](h, zq) |
| 985 | if i_level != 0: |
| 986 | h = self.up[i_level].upsample(h) |
| 987 | |
| 988 | # end |
| 989 | if self.give_pre_end: |
| 990 | return h |
| 991 | |
| 992 | h = self.norm_out(h, zq, clear_fake_cp_cache=clear_fake_cp_cache) |
| 993 | h = nonlinearity(h) |
| 994 | h = self.conv_out(h, clear_cache=clear_fake_cp_cache) |
| 995 | |
| 996 | return h |
| 997 | |
| 998 | def get_last_layer(self): |
| 999 | return self.conv_out.conv.weight |
nothing calls this directly
no test coverage detected