| 26 | |
| 27 | def process_file(lines, rank, args): |
| 28 | def build_input(conversations, tokenizer, args): |
| 29 | zero_width_chars = ["\u200b", "\u200c", "\u200d", "\ufeff"] # filter null characters |
| 30 | for conv in conversations: |
| 31 | if conv['role'] == "assistant": |
| 32 | for char in zero_width_chars: |
| 33 | conv['content'] = conv['content'].replace(char, '') |
| 34 | |
| 35 | if len(conversations) == 0: |
| 36 | return None |
| 37 | |
| 38 | inputs = torch.full((1,), BOS_ID, dtype=torch.int64) |
| 39 | starts = [] |
| 40 | ends = [] |
| 41 | for item in conversations: |
| 42 | content = item["content"] |
| 43 | role = item["role"] |
| 44 | if role == 'system': |
| 45 | cur_inputs = tokenizer(f"<<SYS>>\n{content}\n<</SYS>>\n\n", return_tensors="pt")['input_ids'][0] |
| 46 | elif role == "user": |
| 47 | cur_inputs = tokenizer(f"[INST]{content}[/INST]", return_tensors="pt")['input_ids'][0] |
| 48 | else: |
| 49 | starts.append(inputs.shape[0]) |
| 50 | cur_inputs = tokenizer(content, return_tensors="pt")['input_ids'][0] |
| 51 | ends.append(inputs.shape[0] + cur_inputs.shape[0]) |
| 52 | inputs = torch.cat([inputs, cur_inputs], dim=0) |
| 53 | |
| 54 | inputs = torch.cat([inputs, torch.tensor([EOS_ID])], dim=0) |
| 55 | labels = torch.full_like(inputs, -100) |
| 56 | for start, end in zip(starts, ends): |
| 57 | labels[start:end] = inputs[start:end] |
| 58 | labels[end] = EOS_ID |
| 59 | |
| 60 | if inputs.shape[0] > max_length: |
| 61 | print("exceed_length") |
| 62 | if skip_exceed_length_case: |
| 63 | return None |
| 64 | if truncate_side == 'right': |
| 65 | inputs = inputs[:max_length] |
| 66 | labels = labels[:max_length] |
| 67 | elif truncate_side == 'left': |
| 68 | cut_num = inputs.shape[0] - max_length |
| 69 | inputs = torch.cat([inputs[:2], inputs[2 + cut_num:]], dim=0) |
| 70 | labels = torch.cat([labels[:2], labels[2 + cut_num:]], dim=0) |
| 71 | else: |
| 72 | raise ValueError('truncate_side must be "right" or "left"') |
| 73 | return inputs, labels |
| 74 | |
| 75 | try: |
| 76 | final_inputs = torch.full((len(lines), max_length), PAD_ID, dtype=torch.int64) |