feat: Initial Commit
This commit is contained in:
+113
@@ -0,0 +1,113 @@
|
||||
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
|
||||
Reference in New Issue
Block a user