Evaluate perplexity on the entire dataset
(self, dataset: Dataset)
| 112 | |
| 113 | |
| 114 | def evaluate(self, dataset: Dataset) -> Dict[str, Any]: |
| 115 | """Evaluate perplexity on the entire dataset""" |
| 116 | total_loss = 0 |
| 117 | total_tokens = 0 |
| 118 | total_sentences = 0 |
| 119 | |
| 120 | # Initialize block statistics |
| 121 | num_blocks = (self.block_size - 1) // self.bucket_size + 1 |
| 122 | block_loss = [torch.tensor(0., dtype=torch.float, device=self.device) for _ in range(num_blocks)] |
| 123 | block_tokens = [1e-10 for _ in range(num_blocks)] |
| 124 | bucket_sizes = [0 for _ in range(num_blocks)] |
| 125 | |
| 126 | # Create progress bar |
| 127 | bar = tqdm(self.batchify(dataset, self.block_size)) |
| 128 | |
| 129 | for batch in bar: |
| 130 | batch_outputs = self.process_batch(batch) |
| 131 | input_ids = batch_outputs['input_ids'] |
| 132 | |
| 133 | nlls = batch_outputs['nlls'] |
| 134 | labels = batch_outputs['labels'] |
| 135 | blocks = batch_outputs['blocks'] |
| 136 | |
| 137 | # Update statistics |
| 138 | total_tokens += input_ids.ne(self.loss_fct.ignore_index).sum() |
| 139 | total_sentences += input_ids.shape[0] |
| 140 | print(input_ids.shape[1]) |
| 141 | |
| 142 | for i in blocks: |
| 143 | bucket_sizes[i] += 1 |
| 144 | |
| 145 | # Calculate block-level loss |
| 146 | for i, j in enumerate(range(0, min(input_ids.shape[-1], self.block_size), self.bucket_size)): |
| 147 | block_loss[i] += nlls[:, j:j+self.bucket_size].sum() |
| 148 | block_tokens[i] += labels[:, j:j+self.bucket_size].ne(self.loss_fct.ignore_index).sum() |
| 149 | |
| 150 | # Update total loss |
| 151 | total_loss += batch_outputs['loss'].item() * labels.ne(self.loss_fct.ignore_index).sum() |
| 152 | |
| 153 | # Update progress bar |
| 154 | ppls = [f"{math.exp(loss / toks):6.2f}" for loss, toks in zip(block_loss, block_tokens)] |
| 155 | bar.set_description_str(f"[{total_tokens:10} tokens, {total_sentences:8} sentences] " + ' '.join(ppls)) |
| 156 | |
| 157 | # Calculate final results |
| 158 | final_ppl = math.exp(total_loss / total_tokens) |
| 159 | block_ppls = [math.exp(loss / toks) for loss, toks in zip(block_loss, block_tokens)] |
| 160 | |
| 161 | return { |
| 162 | 'perplexity': final_ppl, |
| 163 | 'block_perplexities': block_ppls, |
| 164 | 'total_tokens': total_tokens, |
| 165 | 'total_sentences': total_sentences |
| 166 | } |
| 167 | |
| 168 | |
| 169 |
no test coverage detected