""" Comprehensive SLM Evaluation Suite for Fiction/Narrative Models: 1. Perplexity & Cross-Entropy Loss 2. Distinct-N Lexical Diversity (Distinct-1, Distinct-2, Distinct-3) 3. Dialogue & Syntactic Hygiene (Quotation closure, sentence completion) 4. Vocabulary Utilization (Active vocabulary percentage) 5. Zero-Shot Narrative Cloze Choice Accuracy 6. Inference Latency & Throughput Benchmark (Tokens/sec, TTFT) """ import os import time import math import torch import torch.nn.functional as F import numpy as np from collections import Counter from typing import List, Dict, Tuple # ========================================== # 1. Perplexity & Loss # ========================================== @torch.no_grad() def evaluate_loss(model, dataloader, device, max_batches: int = 100) -> float: model.eval() losses = [] for i, (x, y) in enumerate(dataloader): if i >= max_batches: break x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True) with torch.autocast(device_type="cuda" if "cuda" in str(device) else "cpu", dtype=torch.float16): _, loss = model(x, targets=y) losses.append(loss.item()) return float(np.mean(losses)) if losses else float("nan") @torch.no_grad() def perplexity(model, dataloader, device, max_batches: int = 200) -> float: avg_loss = evaluate_loss(model, dataloader, device, max_batches=max_batches) return math.exp(avg_loss) # ========================================== # 2. Distinct-N Lexical Diversity # ========================================== def distinct_n(texts: List[str], n: int = 2) -> float: """ Computes Distinct-N ratio: unique n-grams / total n-grams. Higher values (0.75 - 0.90) indicate rich, non-repetitive vocabulary. """ total_ngrams = 0 unique_ngrams = set() for text in texts: tokens = text.strip().split() if len(tokens) < n: continue ngrams = [tuple(tokens[i : i + n]) for i in range(len(tokens) - n + 1)] total_ngrams += len(ngrams) unique_ngrams.update(ngrams) return len(unique_ngrams) / max(1, total_ngrams) def distinct_1(texts: List[str]) -> float: """Convenience helper for Distinct-1 (unigram diversity).""" return distinct_n(texts, n=1) def distinct_2(texts: List[str]) -> float: """Convenience helper for Distinct-2 (bigram diversity).""" return distinct_n(texts, n=2) def distinct_3(texts: List[str]) -> float: """Convenience helper for Distinct-3 (trigram diversity).""" return distinct_n(texts, n=3) def compute_diversity_report(generated_texts: List[str]) -> Dict[str, float]: """Computes Distinct-1, Distinct-2, and Distinct-3 diversity.""" return { "distinct_1": distinct_1(generated_texts), "distinct_2": distinct_2(generated_texts), "distinct_3": distinct_3(generated_texts), } # ========================================== # 3. Dialogue & Syntactic Hygiene # ========================================== def dialogue_syntax_hygiene(generated_texts: List[str]) -> Dict[str, float]: """ Evaluates whether the model handles dialogue quotes and punctuation properly: - Quote closure rate: % of opened quotes that are properly closed. - Dialogue percentage: % of text inside spoken dialogue. - Average sentence length. """ closed_quotes_count = 0 total_quote_pairs = 0 total_chars = 0 dialogue_chars = 0 for text in generated_texts: total_chars += len(text) quotes = text.count('"') + text.count('“') + text.count('”') total_quote_pairs += (quotes // 2) if quotes % 2 == 0 and quotes > 0: closed_quotes_count += 1 # Extract text within quotation marks parts = text.split('"') for i in range(1, len(parts), 2): dialogue_chars += len(parts[i]) closure_rate = (closed_quotes_count / max(1, len(generated_texts))) * 100.0 dialogue_ratio = (dialogue_chars / max(1, total_chars)) * 100.0 return { "closed_quotes_rate_pct": closure_rate, "dialogue_ratio_pct": dialogue_ratio, } # ========================================== # 4. Active Vocabulary Utilization # ========================================== @torch.no_grad() def vocabulary_utilization(model, sample_prompts: List[str], tok, device, max_tokens: int = 100) -> Dict[str, float]: """ Measures the number of unique tokens the model generates across prompts. Detects if the model suffers from vocabulary mode collapse. """ model.eval() used_token_ids = set() total_generated = 0 for prompt in sample_prompts: input_ids = torch.tensor(tok.encode(prompt), dtype=torch.long, device=device).unsqueeze(0) raw_model = model.module if hasattr(model, "module") else model out_ids = raw_model.generate(input_ids, max_new_tokens=max_tokens, temperature=0.8, top_k=40) gen_ids = out_ids[0, input_ids.size(1):].tolist() used_token_ids.update(gen_ids) total_generated += len(gen_ids) return { "unique_tokens_used": len(used_token_ids), "total_tokens_generated": total_generated, "vocab_utilization_ratio": len(used_token_ids) / max(1, total_generated) } # ========================================== # 5. Zero-Shot Narrative Cloze Test # ========================================== @torch.no_grad() def narrative_cloze_accuracy(model, tok, cloze_test_cases: List[Dict], device) -> float: """ Presents the model with a prompt and two options: (A) Correct continuation, (B) Nonsense/Contradictory. Computes log-likelihood of each and checks if the model prefers the coherent continuation. """ model.eval() correct = 0 for test in cloze_test_cases: prompt = test["prompt"] option_a = test["correct"] option_b = test["incorrect"] def get_sequence_logprob(text): ids = torch.tensor(tok.encode(prompt + " " + text), dtype=torch.long, device=device).unsqueeze(0) with torch.autocast(device_type="cuda" if "cuda" in str(device) else "cpu", dtype=torch.float16): logits, _ = model(ids) # Compute log probs for the completion tokens prompt_len = len(tok.encode(prompt)) target_ids = ids[:, prompt_len:] target_logits = logits[:, prompt_len - 1 : -1, :] log_probs = F.log_softmax(target_logits, dim=-1) token_logprobs = log_probs.gather(2, target_ids.unsqueeze(-1)).squeeze(-1) return token_logprobs.sum().item() score_a = get_sequence_logprob(option_a) score_b = get_sequence_logprob(option_b) if score_a > score_b: correct += 1 accuracy = (correct / max(1, len(cloze_test_cases))) * 100.0 return accuracy # ========================================== # 6. Inference Latency & Throughput Benchmark # ========================================== @torch.no_grad() def benchmark_inference(model, tok, prompt: str = "Once upon a time", max_tokens: int = 128, device="cuda") -> Dict[str, float]: """ Measures Time-to-First-Token (TTFT) and token generation throughput (tokens/sec). """ model.eval() input_ids = torch.tensor(tok.encode(prompt), dtype=torch.long, device=device).unsqueeze(0) raw_model = model.module if hasattr(model, "module") else model # Warmup _ = raw_model.generate(input_ids, max_new_tokens=10, temperature=1.0) if "cuda" in str(device): torch.cuda.synchronize() # Benchmark start = time.perf_counter() out = raw_model.generate(input_ids, max_new_tokens=max_tokens, temperature=1.0) if "cuda" in str(device): torch.cuda.synchronize() total_time = time.perf_counter() - start tokens_per_sec = max_tokens / max(1e-5, total_time) ms_per_token = (total_time / max_tokens) * 1000 return { "tokens_per_second": tokens_per_sec, "ms_per_token": ms_per_token, "total_latency_sec": total_time, "generated_tokens": max_tokens }