feat: Initial Commit

This commit is contained in:
2026-09-21 22:20:24 +02:00
commit 7dedbef808
21 changed files with 2457 additions and 0 deletions
+133
View File
@@ -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)
+55
View File
@@ -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)
+196
View File
@@ -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