| import torch |
| from torch.utils.data import Dataset |
| from datasets import load_dataset |
|
|
| from dataset.common import post_processing_chat |
|
|
|
|
| class DPODataset(Dataset): |
| def __init__(self, file_path, tokenizer, max_length=4096): |
| super().__init__() |
| self.tokenizer = tokenizer |
| self.max_length = max_length |
| self.padding = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 0 |
| self.bos_id = tokenizer(f'{tokenizer.bos_token}assistant\n', add_special_tokens=False).input_ids |
| self.eos_id = tokenizer(f'{tokenizer.eos_token}\n', add_special_tokens=False).input_ids |
| self.samples = load_dataset('json', data_files=file_path, split='train') |
|
|
| def __len__(self): |
| return len(self.samples) |
|
|
| def __getitem__(self, index): |
| sample = self.samples[index] |
| chosen = sample['chosen'] |
| rejected = sample['rejected'] |
| chosen_prompt = self.tokenizer.apply_chat_template( |
| chosen, tokenize=False, add_generation_prompt=False |
| ) |
| chosen_prompt = post_processing_chat(chosen_prompt) |
|
|
| rejected_prompt = self.tokenizer.apply_chat_template( |
| rejected, tokenize=False, add_generation_prompt=False |
| ) |
| rejected_prompt = post_processing_chat(rejected_prompt) |
| chosen_encoding = self.tokenizer( |
| chosen_prompt, truncation=True, max_length=self.max_length, padding='max_length' |
| ) |
| rejected_encoding = self.tokenizer( |
| rejected_prompt, truncation=True, max_length=self.max_length, padding='max_length' |
| ) |
|
|
| chosen_input_ids = chosen_encoding['input_ids'] |
| chosen_loss_mask = self.generate_loss_mask(chosen_input_ids) |
|
|
| rejected_input_ids = rejected_encoding['input_ids'] |
| rejected_loss_mask = self.generate_loss_mask(rejected_input_ids) |
| x_chosen = torch.tensor(chosen_input_ids[:-1], dtype=torch.long) |
| y_chosen = torch.tensor(chosen_input_ids[1:], dtype=torch.long) |
| mask_chosen = torch.tensor(chosen_loss_mask[1:], dtype=torch.long) |
| x_rejected = torch.tensor(rejected_input_ids[:-1], dtype=torch.long) |
| y_rejected = torch.tensor(rejected_input_ids[1:], dtype=torch.long) |
| mask_rejected = torch.tensor(rejected_loss_mask[1:], dtype=torch.long) |
|
|
| return { |
| 'x_chosen': x_chosen, |
| 'y_chosen': y_chosen, |
| 'mask_chosen': mask_chosen, |
| 'x_rejected': x_rejected, |
| 'y_rejected': y_rejected, |
| 'mask_rejected': mask_rejected |
| } |
|
|
| def generate_loss_mask(self, input_ids): |
| loss_mask = [0] * len(input_ids) |
| i = 0 |
| while i < len(input_ids): |
| if input_ids[i:i + len(self.bos_id)] == self.bos_id: |
| start = i + len(self.bos_id) |
| end = start |
| while end < len(input_ids): |
| if input_ids[end:end + len(self.eos_id)] == self.eos_id: |
| break |
| end += 1 |
| for j in range(start, min(end + len(self.eos_id), self.max_length)): |
| loss_mask[j] = 1 |
| i = end + len(self.eos_id) if end < len(input_ids) else len(input_ids) |
| else: |
| i += 1 |
| return loss_mask |
|
|