Files
2026-09-21 22:20:24 +02:00

113 lines
2.5 KiB
Python

import json
import torch
from torch.utils.data import Dataset
from torch.nn.utils.rnn import pad_sequence
from config import (
PAD_TOKEN,
MAX_PROMPT_LENGTH,
MAX_OUTPUT_LENGTH,
)
class MinecraftDataset(Dataset):
"""
Dataset for training the Minecraft Builder AI.
"""
def __init__(self, dataset_path, tokenizer):
self.tokenizer = tokenizer
with open(dataset_path, "r", encoding="utf8") as f:
self.data = json.load(f)
self.pad_text = tokenizer.text_to_id[PAD_TOKEN]
self.pad_output = tokenizer.output_to_id[PAD_TOKEN]
def __len__(self):
return len(self.data)
def __getitem__(self, index):
sample = self.data[index]
prompt_ids = self.tokenizer.encode_prompt(
sample["prompt"]
)
output_ids = self.tokenizer.encode_blocks(
sample["blocks"]
)
# Limit maximum sequence lengths
prompt_ids = prompt_ids[:MAX_PROMPT_LENGTH]
output_ids = output_ids[:MAX_OUTPUT_LENGTH]
return (
torch.tensor(prompt_ids, dtype=torch.long),
torch.tensor(output_ids, dtype=torch.long),
)
def collate_fn(batch):
"""
Pads sequences inside a batch.
"""
prompts = [item[0] for item in batch]
outputs = [item[1] for item in batch]
prompt_pad = prompts[0].new_tensor(
[0]
) # placeholder (overwritten below)
output_pad = outputs[0].new_tensor(
[0]
)
# These values are replaced by the DataLoader factory below.
prompt_padding_value = getattr(collate_fn, "prompt_pad", 0)
output_padding_value = getattr(collate_fn, "output_pad", 0)
prompts = pad_sequence(
prompts,
batch_first=True,
padding_value=prompt_padding_value,
)
outputs = pad_sequence(
outputs,
batch_first=True,
padding_value=output_padding_value,
)
return prompts, outputs
def create_dataloader(
dataset_path,
tokenizer,
batch_size,
shuffle=True,
):
"""
Creates a DataLoader with automatic padding.
"""
dataset = MinecraftDataset(
dataset_path,
tokenizer,
)
# Give the collate function the correct padding IDs
collate_fn.prompt_pad = tokenizer.text_to_id[PAD_TOKEN]
collate_fn.output_pad = tokenizer.output_to_id[PAD_TOKEN]
loader = torch.utils.data.DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
collate_fn=collate_fn,
)
return loader