import torch from torch.utils.data import Dataset from accelerate import Accelerator, DistributedType class LineByLineTextDataset(Dataset): def __init__(self, tokenizer, raw_datasets, max_length: int): self.padding = "max_length" self.text_column_name = 'text' self.max_length = max_length self.accelerator = Accelerator(gradient_accumulation_steps=1) self.tokenizer = tokenizer with self.accelerator.main_process_first(): self.tokenized_datasets = raw_datasets.map( self.tokenize_function, batched=True, num_proc=4, remove_columns=[self.text_column_name], desc="Running tokenizer on dataset line_by_line", ) self.tokenized_datasets.set_format('torch',columns=['input_ids'],dtype=torch.long) def tokenize_function(self,examples): examples[self.text_column_name] = [ line for line in examples[self.text_column_name] if len(line[0]) > 0 and not line[0].isspace() ] return self.tokenizer( examples[self.text_column_name], padding=self.padding, truncation=True, max_length=self.max_length, return_special_tokens_mask=True, ) def __len__(self): return len(self.tokenized_datasets) def __getitem__(self, i): return self.tokenized_datasets[i]