Download repetition_diagnostic.py from Compactbot/repetition-diagnostic: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/Compactbot/repetition-diagnostic/resolve/main/repetition_diagnostic.py
- Command line
-
hf download hf://Compactbot/repetition-diagnostic/repetition_diagnostic.py
-
curl -L -o repetition_diagnostic.py https://huggingface.co/Compactbot/repetition-diagnostic/resolve/main/repetition_diagnostic.py
10.2 kB
| #!/usr/bin/env python3 | |
| """ | |
| repetition_probe.py — A builder tool for tiny model developers. | |
| Takes a HuggingFace model and probes how different repetition_penalty values | |
| affect generation quality. Outputs a comparison table with: | |
| - Repetition rate (fraction of tokens already seen in the output) | |
| - Unique token ratio (diversity) | |
| - Loop detection (repeated n-grams) | |
| - Sample text | |
| Usage: | |
| python3 repetition_probe.py --model kevin-bretz/Small-Language-Model --prompt "Once upon a time" | |
| python3 repetition_probe.py --model kevin-bretz/Small-Language-Model --prompts "Once upon a time|The cat sat|I want to" | |
| Why this matters: | |
| Tiny models (<50M params) are extremely prone to degenerate repetition loops. | |
| The default repetition_penalty=1.0 means NO penalty — the model can output the | |
| same token forever. Most training scripts don't test generation quality, so | |
| you discover this at inference time with no idea what to do. | |
| This tool gives you a quick diagnostic: sweep the penalty, see where quality | |
| peaks, and know what to set in your inference config. | |
| """ | |
| import argparse | |
| import json | |
| import sys | |
| import time | |
| from collections import Counter | |
| from dataclasses import dataclass, field | |
| from typing import List, Optional | |
| import torch | |
| import torch.nn.functional as F | |
| class ProbeResult: | |
| prompt: str | |
| penalty: float | |
| text: str | |
| n_tokens: int | |
| repetition_rate: float # fraction of output tokens that appeared before | |
| unique_token_ratio: float # unique tokens / total tokens | |
| max_loop_len: int # longest repeated trigram | |
| n_repeated_trigrams: int # count of trigrams appearing 2+ times | |
| elapsed_s: float | |
| def compute_repetition_rate(token_ids: List[int]) -> float: | |
| """Fraction of tokens (excluding first) that have appeared earlier in the sequence.""" | |
| if len(token_ids) < 2: | |
| return 0.0 | |
| seen = set() | |
| repeated = 0 | |
| for t in token_ids[1:]: | |
| if t in seen: | |
| repeated += 1 | |
| seen.add(t) | |
| return repeated / (len(token_ids) - 1) | |
| def compute_unique_ratio(token_ids: List[int]) -> float: | |
| """Unique tokens / total tokens.""" | |
| if not token_ids: | |
| return 0.0 | |
| return len(set(token_ids)) / len(token_ids) | |
| def find_repeated_trigrams(token_ids: List[int]) -> tuple: | |
| """Find the longest repeated trigram and count of repeated trigrams.""" | |
| if len(token_ids) < 6: | |
| return 0, 0 | |
| trigrams = [] | |
| for i in range(len(token_ids) - 2): | |
| trigrams.append(tuple(token_ids[i:i+3])) | |
| counts = Counter(trigrams) | |
| repeated = sum(1 for c in counts.values() if c >= 2) | |
| # Longest repeated trigram: find the longest n-gram that repeats | |
| max_len = 0 | |
| for n in range(3, min(len(token_ids), 20)): | |
| ngrams = [tuple(token_ids[i:i+n]) for i in range(len(token_ids) - n + 1)] | |
| ngram_counts = Counter(ngrams) | |
| if any(c >= 2 for c in ngram_counts.values()): | |
| max_len = n | |
| else: | |
| break | |
| return max_len, repeated | |
| def generate_with_penalty( | |
| model, | |
| tokenizer, | |
| prompt: str, | |
| max_new_tokens: int = 64, | |
| repetition_penalty: float = 1.0, | |
| temperature: float = 0.8, | |
| top_k: int = 50, | |
| device: str = "cpu", | |
| ) -> tuple: | |
| """Generate text with a given repetition penalty. Returns (text, token_ids, elapsed).""" | |
| model.eval() | |
| input_ids = tokenizer.encode(prompt, return_tensors="pt").to(device) | |
| gen_start = input_ids.shape[1] | |
| with torch.no_grad(): | |
| generated = input_ids | |
| for _ in range(max_new_tokens): | |
| outputs = model(generated) | |
| logits = outputs.logits[0, -1, :] | |
| # Apply repetition penalty | |
| if repetition_penalty != 1.0: | |
| for i in range(generated.shape[1]): | |
| tok = generated[0, i].item() | |
| if logits[tok] > 0: | |
| logits[tok] /= repetition_penalty | |
| else: | |
| logits[tok] *= repetition_penalty | |
| # Temperature | |
| if temperature != 1.0: | |
| logits = logits / temperature | |
| # Top-k | |
| if top_k > 0: | |
| top_k_val = min(top_k, logits.shape[-1]) | |
| indices_to_remove = logits < torch.topk(logits, top_k_val)[0][..., -1:] | |
| logits[indices_to_remove] = float("-inf") | |
| probs = F.softmax(logits, dim=-1) | |
| next_token = torch.multinomial(probs, num_samples=1).unsqueeze(0) | |
| generated = torch.cat([generated, next_token], dim=1) | |
| # Stop on EOS | |
| if next_token.item() == tokenizer.eos_token_id: | |
| break | |
| elapsed = None # measured outside | |
| output_ids = generated[0, gen_start:].tolist() | |
| text = tokenizer.decode(output_ids, skip_special_tokens=True) | |
| return text, output_ids | |
| def run_probe( | |
| model_name: str, | |
| prompts: List[str], | |
| penalties: List[float], | |
| max_new_tokens: int = 64, | |
| temperature: float = 0.8, | |
| top_k: int = 50, | |
| device: str = "cpu", | |
| seed: int = 42, | |
| ) -> List[ProbeResult]: | |
| """Run the full probe sweep.""" | |
| print(f"Loading model: {model_name}") | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| tokenizer = AutoTokenizer.from_pretrained(model_name) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=torch.float32).to(device) | |
| torch.manual_seed(seed) | |
| results = [] | |
| for prompt in prompts: | |
| for penalty in penalties: | |
| torch.manual_seed(seed) # Same seed for fair comparison | |
| start = time.time() | |
| text, token_ids = generate_with_penalty( | |
| model, tokenizer, prompt, | |
| max_new_tokens=max_new_tokens, | |
| repetition_penalty=penalty, | |
| temperature=temperature, | |
| top_k=top_k, | |
| device=device, | |
| ) | |
| elapsed = time.time() - start | |
| rep_rate = compute_repetition_rate(token_ids) | |
| unique_ratio = compute_unique_ratio(token_ids) | |
| max_loop, n_rep_tri = find_repeated_trigrams(token_ids) | |
| results.append(ProbeResult( | |
| prompt=prompt, | |
| penalty=penalty, | |
| text=text, | |
| n_tokens=len(token_ids), | |
| repetition_rate=rep_rate, | |
| unique_token_ratio=unique_ratio, | |
| max_loop_len=max_loop, | |
| n_repeated_trigrams=n_rep_tri, | |
| elapsed_s=elapsed, | |
| )) | |
| print(f" [{prompt[:20]:20s}] rp={penalty:.1f} rep={rep_rate:.1%} " | |
| f"uniq={unique_ratio:.1%} max_loop={max_loop} " | |
| f"n_rep_tri={n_rep_tri} ({len(token_ids)} tokens)") | |
| return results | |
| def format_table(results: List[ProbeResult]) -> str: | |
| """Format results as a readable table.""" | |
| lines = [] | |
| lines.append(f"{'Prompt':<22} {'RP':>4} {'Tokens':>6} {'Repet':>7} {'Unique':>7} {'MaxLoop':>8} {'RepTri':>7}") | |
| lines.append("-" * 72) | |
| for r in results: | |
| lines.append( | |
| f"{r.prompt[:20]:<22} {r.penalty:>4.1f} {r.n_tokens:>6} " | |
| f"{r.repetition_rate:>7.1%} {r.unique_token_ratio:>7.1%} " | |
| f"{r.max_loop_len:>8} {r.n_repeated_trigrams:>7}" | |
| ) | |
| lines.append("-" * 72) | |
| lines.append("") | |
| lines.append("Sample outputs (first 100 chars):") | |
| lines.append("") | |
| for r in results: | |
| lines.append(f" [{r.prompt[:15]:15s}] rp={r.penalty:.1f}: {r.text[:100]}") | |
| return "\n".join(lines) | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Repetition penalty probe for tiny models") | |
| parser.add_argument("--model", required=True, help="HuggingFace model name or path") | |
| parser.add_argument("--prompt", default="Once upon a time", help="Single prompt") | |
| parser.add_argument("--prompts", default=None, help="Pipe-separated prompts") | |
| parser.add_argument("--penalties", default="1.0,1.1,1.2,1.3,1.5", | |
| help="Comma-separated repetition penalty values") | |
| parser.add_argument("--max-tokens", type=int, default=64, help="Max new tokens to generate") | |
| parser.add_argument("--temperature", type=float, default=0.8) | |
| parser.add_argument("--top-k", type=int, default=50) | |
| parser.add_argument("--device", default="cpu") | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--output", default=None, help="Save JSON results to file") | |
| args = parser.parse_args() | |
| if args.prompts: | |
| prompts = [p.strip() for p in args.prompts.split("|")] | |
| else: | |
| prompts = [args.prompt] | |
| penalties = [float(p) for p in args.penalties.split(",")] | |
| print(f"Model: {args.model}") | |
| print(f"Prompts: {prompts}") | |
| print(f"Penalties: {penalties}") | |
| print(f"Max tokens: {args.max_tokens}, Temp: {args.temperature}, Top-k: {args.top_k}") | |
| print() | |
| results = run_probe( | |
| args.model, prompts, penalties, | |
| max_new_tokens=args.max_tokens, | |
| temperature=args.temperature, | |
| top_k=args.top_k, | |
| device=args.device, | |
| seed=args.seed, | |
| ) | |
| print() | |
| print(format_table(results)) | |
| if args.output: | |
| out = { | |
| "model": args.model, | |
| "prompts": prompts, | |
| "penalties": penalties, | |
| "max_tokens": args.max_tokens, | |
| "temperature": args.temperature, | |
| "top_k": args.top_k, | |
| "seed": args.seed, | |
| "results": [ | |
| { | |
| "prompt": r.prompt, | |
| "penalty": r.penalty, | |
| "text": r.text, | |
| "n_tokens": r.n_tokens, | |
| "repetition_rate": round(r.repetition_rate, 4), | |
| "unique_token_ratio": round(r.unique_token_ratio, 4), | |
| "max_loop_len": r.max_loop_len, | |
| "n_repeated_trigrams": r.n_repeated_trigrams, | |
| } | |
| for r in results | |
| ], | |
| } | |
| with open(args.output, "w") as f: | |
| json.dump(out, f, indent=2) | |
| print(f"\nResults saved to {args.output}") | |
| if __name__ == "__main__": | |
| main() |