MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / forward

Method forward

sat/vae_modules/cp_enc_dec.py:963–996  ·  view source on GitHub ↗
(self, z, clear_fake_cp_cache=True)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

nonlinearityFunction · 0.70

Tested by

no test coverage detected