113 lines
2.5 KiB
Python
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 |