MCPcopy Create free account
hub / github.com/coperception/star / forward_decoder

Method forward_decoder

star/models/VQSTAR.py:359–392  ·  view source on GitHub ↗

overwrite the original forward_decoder, now input latent are fused

(self, latent, mask, ids_restore)

Source from the content-addressed store, hash-verified

357 return x, mask, ids_restore
358
359 def forward_decoder(self, latent, mask, ids_restore):
360 """
361 overwrite the original forward_decoder, now input latent are fused
362 """
363 # print("to decoder, latent", latent.size())
364 latent = self.decompressor(latent)
365 # # embed tokens
366 latent = self.decoder_embed(latent)
367
368 restored_latent = self.unmasking_handle(latent, mask, ids_restore)
369 restored_latent = restored_latent.reshape(restored_latent.size(0), self.patch_h*self.patch_w, -1)
370 # x = torch.cat([latent_cls, restored_latent], dim=1)
371
372 # add pos embed
373 x = restored_latent + self.decoder_pos_embed[:, 1:, :]
374
375 # apply Transformer blocks
376 for blk in self.decoder_blocks:
377 x = blk(x)
378 x = self.decoder_norm(x)
379
380 # VERSION mlp -----
381 x_occ = self.decoder_pred_occ(x)
382 x_free = self.decoder_pred_free(x)
383 x_occ = self.unpatchify(x_occ)
384 x_free = self.unpatchify(x_free)
385 # remove cls token
386 # x = x[:, 1:, :]
387 # print("pred size", x.size())
388 # --------------------
389 # x_pred = x_occ
390 # print("x_pred", x_pred.size())
391 x_pred = torch.stack((x_free, x_occ), dim=1) # [B, class, C, H, W]
392 return x_pred
393
394 def forward_loss(self, teacher, pred):
395 """

Callers 2

forwardMethod · 0.95
inferenceMethod · 0.95

Calls 1

unpatchifyMethod · 0.80

Tested by

no test coverage detected