196 lines
4.5 KiB
Python
196 lines
4.5 KiB
Python
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 |