feat: Initial Commit
This commit is contained in:
@@ -0,0 +1,133 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from config import (
|
||||
DEVICE,
|
||||
START_TOKEN,
|
||||
END_TOKEN,
|
||||
MAX_OUTPUT_LENGTH,
|
||||
)
|
||||
|
||||
from model.transformer import MinecraftTransformer
|
||||
|
||||
|
||||
class BuilderModel(nn.Module):
|
||||
"""
|
||||
High-level wrapper around the MinecraftTransformer.
|
||||
"""
|
||||
|
||||
def __init__(self, tokenizer):
|
||||
super().__init__()
|
||||
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.transformer = MinecraftTransformer(
|
||||
tokenizer.text_vocab_size,
|
||||
tokenizer.output_vocab_size,
|
||||
)
|
||||
|
||||
def forward(self, src, tgt):
|
||||
"""
|
||||
Training forward pass.
|
||||
|
||||
src = prompt tokens
|
||||
tgt = decoder input tokens (already shifted by train.py)
|
||||
"""
|
||||
|
||||
return self.transformer(
|
||||
src=src,
|
||||
tgt=tgt,
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(
|
||||
self,
|
||||
prompt,
|
||||
max_length=MAX_OUTPUT_LENGTH,
|
||||
):
|
||||
"""
|
||||
Generate a Minecraft structure from a prompt.
|
||||
"""
|
||||
|
||||
self.eval()
|
||||
|
||||
device = next(self.parameters()).device
|
||||
|
||||
# Encode prompt
|
||||
src = torch.tensor(
|
||||
[self.tokenizer.encode_prompt(prompt)],
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# Encode with transformer
|
||||
memory = self.transformer.encode(src)
|
||||
|
||||
start_id = self.tokenizer.output_to_id[START_TOKEN]
|
||||
end_id = self.tokenizer.output_to_id[END_TOKEN]
|
||||
|
||||
generated = [start_id]
|
||||
|
||||
for _ in range(max_length):
|
||||
|
||||
tgt = torch.tensor(
|
||||
[generated],
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
logits = self.transformer.decode(
|
||||
tgt,
|
||||
memory,
|
||||
)
|
||||
|
||||
probs = torch.softmax(logits[0, -1], dim=-1)
|
||||
|
||||
next_token = torch.argmax(probs).item()
|
||||
|
||||
print(
|
||||
next_token,
|
||||
self.tokenizer.id_to_output[next_token],
|
||||
f"{probs[next_token].item():.3f}"
|
||||
)
|
||||
|
||||
generated.append(next_token)
|
||||
|
||||
if next_token == end_id:
|
||||
break
|
||||
|
||||
blocks = self.tokenizer.decode_blocks(generated)
|
||||
|
||||
# Don't count the <START> token
|
||||
token_count = max(0, len(generated) - 1)
|
||||
|
||||
return blocks, token_count
|
||||
|
||||
def save(self, path):
|
||||
"""
|
||||
Save model weights.
|
||||
"""
|
||||
|
||||
torch.save(
|
||||
self.state_dict(),
|
||||
path,
|
||||
)
|
||||
|
||||
def load(self, path):
|
||||
"""
|
||||
Load model weights.
|
||||
"""
|
||||
|
||||
self.load_state_dict(
|
||||
torch.load(
|
||||
path,
|
||||
map_location=DEVICE,
|
||||
)
|
||||
)
|
||||
|
||||
def to_device(self):
|
||||
"""
|
||||
Move model to configured device.
|
||||
"""
|
||||
|
||||
return self.to(DEVICE)
|
||||
@@ -0,0 +1,55 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class PositionalEncoding(nn.Module):
|
||||
"""
|
||||
Standard sinusoidal positional encoding from
|
||||
"Attention Is All You Need".
|
||||
"""
|
||||
|
||||
def __init__(self, embed_dim, dropout=0.1, max_len=10000):
|
||||
super().__init__()
|
||||
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
pe = torch.zeros(max_len, embed_dim)
|
||||
|
||||
position = torch.arange(
|
||||
0,
|
||||
max_len,
|
||||
dtype=torch.float
|
||||
).unsqueeze(1)
|
||||
|
||||
div_term = torch.exp(
|
||||
torch.arange(
|
||||
0,
|
||||
embed_dim,
|
||||
2
|
||||
).float() *
|
||||
(-math.log(10000.0) / embed_dim)
|
||||
)
|
||||
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
|
||||
pe = pe.unsqueeze(0)
|
||||
|
||||
# Stored as a buffer so it moves with the model
|
||||
# but isn't trained.
|
||||
self.register_buffer("pe", pe)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Args:
|
||||
x: Tensor of shape
|
||||
(batch_size, sequence_length, embed_dim)
|
||||
"""
|
||||
|
||||
seq_len = x.size(1)
|
||||
|
||||
x = x + self.pe[:, :seq_len]
|
||||
|
||||
return self.dropout(x)
|
||||
@@ -0,0 +1,196 @@
|
||||
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
|
||||
Reference in New Issue
Block a user