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

Function main

evals/ppl.py:170–225  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

168
169
170def main():
171 parser = argparse.ArgumentParser(description="Evaluate perplexity")
172 parser.add_argument('-p', '--path', type=str, default='fla-hub/gla-1.3B-100B')
173 parser.add_argument('-d', '--data', type=str, default='fla-hub/slimpajama-test')
174 parser.add_argument('-s', '--split', type=str, default='train')
175 parser.add_argument('-n', '--column_name', type=str, default='text')
176 parser.add_argument('--block_size', type=int, default=28672)
177 parser.add_argument('--bucket_size', type=int, default=2048)
178 parser.add_argument('--batch_size', type=int, default=1)
179 args = parser.parse_args()
180
181 # Set device and random seed
182 device = "cuda"
183 torch.manual_seed(0)
184
185 # Load model and tokenizer
186 print(f"Loading model {args.path}")
187 tokenizer = AutoTokenizer.from_pretrained(args.path)
188 model = AutoModelForCausalLM.from_pretrained(
189 args.path,
190 device_map={"": device}
191 ).bfloat16().eval()
192 print(f"{model}")
193
194 # Load dataset
195 print(f"Loading data {args.data}")
196 dataset = load_dataset(args.data, split=args.split)
197 dataset = dataset.map(
198 partial(PerplexityEvaluator.preprocess, tokenizer=tokenizer, column_name=args.column_name),
199 batched=True,
200 num_proc=32
201 )
202 print(dataset)
203 print("batch_size", args.batch_size, "block_size", args.block_size, "total_tokens_per_batch", args.batch_size * args.block_size)
204
205 # Create evaluator and run evaluation
206 evaluator = PerplexityEvaluator(
207 model=model,
208 tokenizer=tokenizer,
209 device=device,
210 block_size=args.block_size,
211 bucket_size=args.bucket_size,
212 batch_size=args.batch_size
213 )
214
215 with torch.no_grad():
216 results = evaluator.evaluate(dataset)
217
218 # Print results
219 print("\nEvaluation Results:")
220 print(f"Final Perplexity: {results['perplexity']:.2f}")
221 print(f"Total Tokens: {results['total_tokens']}")
222 print(f"Total Sentences: {results['total_sentences']}")
223 print("\nBlock-wise Perplexities:")
224 for i, ppl in enumerate(results['block_perplexities']):
225 print(f"Block {i}: {ppl:.2f}")
226
227if __name__ == "__main__":

Callers 1

ppl.pyFile · 0.70

Calls 2

evaluateMethod · 0.95
PerplexityEvaluatorClass · 0.85

Tested by

no test coverage detected