#!/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 @dataclass 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()