JavRedstone's picture
download
raw
1.97 kB
import torch
from torch.utils.data import Dataset
from torch.nn.utils.rnn import pad_sequence
from typing import List, Tuple
from data.indirect_idx.tokenizer import CharacterTokenizer
class IndirectIdxDataset(Dataset):
def __init__(self, data: List[str], tokenizer: CharacterTokenizer):
self.data = data
self.tokenizer = tokenizer
self.tokenized_data = []
# tokenize all data once
for sample in data:
parts = sample.split(', ')
source_str, source_char, shift_str, target_char = parts
delimiter = ", "
input_str = delimiter.join([source_str, source_char, shift_str])
input_tokens = self.tokenizer.encode_input(input_str)
target_token = self.tokenizer.encode_target(target_char.replace('\n', ''))
self.tokenized_data.append((input_tokens, target_token))
def __len__(self) -> int:
return len(self.tokenized_data)
def __getitem__(self, idx: int) -> Tuple[List[int], int]:
return self.tokenized_data[idx]
def collate_fn(batch: List[Tuple[List[int], int]], pad_idx: int = 0):
inputs, targets = zip(*batch)
input_tensors = [torch.tensor(seq, dtype=torch.long) for seq in inputs]
padded_inputs = pad_sequence(input_tensors, batch_first=True, padding_value=pad_idx)
target_tensor = torch.full_like(padded_inputs, fill_value=-1, dtype=torch.long)
target_masks = []
for i, (padded_input, input_tensor) in enumerate(zip(padded_inputs, input_tensors)):
mask = torch.zeros_like(padded_input, dtype=torch.bool)
indices = torch.where(padded_input == input_tensor[-1])[0]
last_idx = indices[-1]
mask[last_idx] = True
target_tensor[i, last_idx] = targets[i]
target_masks.append(mask)
# attention_mask = (padded_inputs != pad_idx).long()
target_masks = torch.stack(target_masks, dim=0)
return padded_inputs, target_tensor, target_masks

Xet Storage Details

Size:
1.97 kB
·
Xet hash:
5739c39b1d73bcf8fe30dce8e7e22e806b25fdbabec5f62b3ae35611c8ae0706

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.