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