MCPcopy Create free account
hub / github.com/OpenSparseLLMs/MoM / evaluate

Method evaluate

evals/ppl.py:114–166  ·  view source on GitHub ↗

Evaluate perplexity on the entire dataset

(self, dataset: Dataset)

Source from the content-addressed store, hash-verified

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

Callers 1

mainFunction · 0.95

Calls 2

batchifyMethod · 0.95
process_batchMethod · 0.95

Tested by

no test coverage detected