repetition-diagnostic / repetition_diagnostic.py
Compactbot's picture
Add repetition penalty diagnostic script (#1)
a80c66a
Raw History Blame Contribute Delete
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
@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()