| 75 | self.tokenizer = tokenizer |
| 76 | |
| 77 | def __call__(self, example, has_aux=False, add_bos_token=True, add_eos_token=True): |
| 78 | if has_aux: |
| 79 | example, *aux = example |
| 80 | else: |
| 81 | aux = tuple() |
| 82 | token_buffer = [] |
| 83 | loss_mask_buffer = [] |
| 84 | |
| 85 | if add_bos_token and self.config.add_bos_token: |
| 86 | token_buffer.append(self.tokenizer.bos_token_id) |
| 87 | loss_mask_buffer.append(0.0) |
| 88 | |
| 89 | if self.config.fields_from_example != '': |
| 90 | fields = example[self.config.fields_from_example].split(',') |
| 91 | else: |
| 92 | fields = self.config.fields.split(',') |
| 93 | |
| 94 | for i, field in enumerate(fields): |
| 95 | if field.startswith('[') and field.endswith(']'): |
| 96 | # No loss for this field. |
| 97 | field = field[1:-1] |
| 98 | mask = 0.0 |
| 99 | else: |
| 100 | mask = 1.0 |
| 101 | |
| 102 | if field == '<|bos|>': |
| 103 | token_buffer.append(self.tokenizer.bos_token_id) |
| 104 | loss_mask_buffer.append(mask) |
| 105 | elif field == '<|eos|>': |
| 106 | token_buffer.append(self.tokenizer.eos_token_id) |
| 107 | loss_mask_buffer.append(mask) |
| 108 | else: |
| 109 | subfields = field.split('+') |
| 110 | text = self.config.subfield_separator.join( |
| 111 | [example[subfield] for subfield in subfields] |
| 112 | ) |
| 113 | if i == 0: |
| 114 | text = self.config.prepend_text + text |
| 115 | tokens = self.tokenizer.encode(text, add_special_tokens=False) |
| 116 | token_buffer.extend(tokens) |
| 117 | loss_mask_buffer.extend([mask for _ in range(len(tokens))]) |
| 118 | |
| 119 | if add_eos_token and self.config.add_eos_token: |
| 120 | token_buffer.append(self.tokenizer.eos_token_id) |
| 121 | loss_mask_buffer.append(1.0) |
| 122 | |
| 123 | return token_buffer, loss_mask_buffer, *aux |
| 124 | |
| 125 | |
| 126 | class VisionTextProcessor(object): |