def calculate_perplexity(model, text): # Encode the text encodings = tokenizer(text, return_tensors='pt').to(device) # Define input_ids and target_ids input_ids = encodings.input_ids target_ids = input_ids.clone() with torch.no_grad(): outputs = model(input_ids, labels=target_ids) # Loss calculation neg_log_likelihood = outputs.loss # Perplexity calculation ppl = torch.exp(neg_log_likelihood) return ppl ppl = calculate_perplexity(model, original_text) ppl_abs = calculate_perplexity(model_abs, absmax_text) ppl_zp = calculate_perplexity(model_zp, absmax_text) print(f"Original perplexity: {ppl.item():.2f}") print(f"Absmax perplexity: {ppl_abs.item():.2f}") print(f"Zeropoint perplexity: {ppl_zp.item():.2f}")