Views
No views yet
1from compressor.decoder import PerceiverDecoder
2
3model = PerceiverDecoder(input_dim=512, output_dim=1664, num_queries=576)
4model.load_state_dict(torch.load("decoder_stepX_hrsY.pt"))
5# Input: [B, 64, 512] Perceiver output → Output: [B, 576, 1664] V-JEPA latents