import torch import torch.nn as nn from config import ( EMBED_DIM, NUM_HEADS, NUM_ENCODER_LAYERS, NUM_DECODER_LAYERS, FEED_FORWARD_DIM, DROPOUT, ) from model.positional_encoding import PositionalEncoding class MinecraftTransformer(nn.Module): """ Encoder-Decoder Transformer used by the Minecraft Builder AI. """ def __init__( self, text_vocab_size, output_vocab_size, ): super().__init__() self.embed_dim = EMBED_DIM # ---------------------------- # Embeddings # ---------------------------- self.text_embedding = nn.Embedding( text_vocab_size, EMBED_DIM ) self.output_embedding = nn.Embedding( output_vocab_size, EMBED_DIM ) # ---------------------------- # Positional Encoding # ---------------------------- self.text_position = PositionalEncoding( EMBED_DIM, DROPOUT ) self.output_position = PositionalEncoding( EMBED_DIM, DROPOUT ) # ---------------------------- # Transformer # ---------------------------- self.transformer = nn.Transformer( d_model=EMBED_DIM, nhead=NUM_HEADS, num_encoder_layers=NUM_ENCODER_LAYERS, num_decoder_layers=NUM_DECODER_LAYERS, dim_feedforward=FEED_FORWARD_DIM, dropout=DROPOUT, batch_first=True, ) # ---------------------------- # Output layer # ---------------------------- self.fc_out = nn.Linear( EMBED_DIM, output_vocab_size ) # ================================================== # Masks # ================================================== def generate_square_subsequent_mask(self, size, device): """ Prevent the decoder from seeing future tokens. """ return torch.triu( torch.full( (size, size), float("-inf"), device=device ), diagonal=1 ) # ================================================== # Forward # ================================================== def forward( self, src, tgt, src_padding_mask=None, tgt_padding_mask=None, ): """ Parameters ---------- src : (batch, src_len) tgt : (batch, tgt_len) Returns ------- logits : (batch, tgt_len, output_vocab_size) """ src = self.text_embedding(src) tgt = self.output_embedding(tgt) src = self.text_position(src) tgt = self.output_position(tgt) tgt_mask = self.generate_square_subsequent_mask( tgt.size(1), tgt.device ) output = self.transformer( src=src, tgt=tgt, tgt_mask=tgt_mask, src_key_padding_mask=src_padding_mask, tgt_key_padding_mask=tgt_padding_mask, memory_key_padding_mask=src_padding_mask, ) logits = self.fc_out(output) return logits # ================================================== # Encoder # ================================================== def encode( self, src, src_padding_mask=None, ): src = self.text_embedding(src) src = self.text_position(src) memory = self.transformer.encoder( src, src_key_padding_mask=src_padding_mask ) return memory # ================================================== # Decoder # ================================================== def decode( self, tgt, memory, tgt_padding_mask=None, memory_padding_mask=None, ): tgt = self.output_embedding(tgt) tgt = self.output_position(tgt) tgt_mask = self.generate_square_subsequent_mask( tgt.size(1), tgt.device ) output = self.transformer.decoder( tgt=tgt, memory=memory, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_padding_mask, memory_key_padding_mask=memory_padding_mask, ) logits = self.fc_out(output) return logits