# Codeblock 21 class DecoderTorch(nn.Module): def __init__(self): super().__init__() self.embedding = nn.Embedding(num_embeddings=VOCAB_SIZE, embedding_dim=EMBED_DIM) self.sinusoidal_embedding = SinusoidalEmbedding() #(1) decoder_block = nn.TransformerDecoderLayer(d_model=EMBED_DIM, nhead=NUM_HEADS, dim_feedforward=HIDDEN_DIM, dropout=DROP_PROB, batch_first=True) #(2) self.decoder_blocks = nn.TransformerDecoder(decoder_layer=decoder_block, num_layers=NUM_DECODER_BLOCKS) self.linear = nn.Linear(in_features=EMBED_DIM, out_features=VOCAB_SIZE) def forward(self, features, captions, tgt_mask): print(f"features\t\t: {features.shape}") print(f"captions\t\t: {captions.shape}") captions = self.embedding(captions) print(f"after embedding\t\t: {captions.shape}") captions = captions + self.sinusoidal_embedding() print(f"after sin embed\t\t: {captions.shape}") #(3) captions = self.decoder_blocks(tgt=captions, memory=features, tgt_mask=tgt_mask) print(f"after decoder blocks\t: {captions.shape}") captions = self.linear(captions) print(f"after linear\t\t: {captions.shape}") return captions