#!/usr/bin/env python """Run the three controlled OLMo 2 1B conditions used in the RWKV comparison.""" import argparse import csv import json import math from collections import Counter from datetime import date from pathlib import Path import torch from transformers import AutoModelForCausalLM, AutoTokenizer MODEL_ID = "allenai/OLMo-2-0425-1B-Instruct" USER_PROMPT = ( "정수 a,b가 ab+1 | a^2+b^2를 만족할 때, " "(a^2+b^2)/(ab+1)이 완전제곱수임을 증명해. 단, a,b는 양의 정수다." ) SYSTEM_GIVE_UP = ( "문제를 풀 수 없거나 답을 모르겠으면, 억지로 계속 추측하지 말고 " "“모르겠습니다”라고 답한 뒤 멈춰라." ) def logits_metrics(logits): logits = logits.float() log_probs = torch.log_softmax(logits, dim=-1) probs = log_probs.exp() top = torch.topk(logits, k=2) return { "top1_probability": probs[top.indices[0]], "entropy_nats": -(probs * log_probs).sum(), "top1_top2_margin": top.values[0] - top.values[1], "top1_token_id": top.indices[0], } def rolling_dup8_rate(tokens, window=64, n=8): tail = tokens[-window:] grams = [tuple(tail[i : i + n]) for i in range(max(0, len(tail) - n + 1))] if not grams: return 0.0 return (len(grams) - len(set(grams))) / len(grams) def largest_repeated_ngram(tokens): """Largest exact repeated span with non-overlapping occurrences.""" count = len(tokens) for n in range(count // 2, 0, -1): starts = {} for i in range(count - n + 1): gram = tuple(tokens[i : i + n]) earlier = starts.setdefault(gram, []) for j in earlier: if i >= j + n: return n, j + 1, i + 1 earlier.append(i) return 0, None, None def run_one(model, tokenizer, condition, messages, repetition_penalty, max_new_tokens, out_dir): encoded = tokenizer.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to("cuda") prompt_len = encoded["input_ids"].shape[-1] raw_logits_rows = [] def capture_raw_logits(_module, _inputs, output): if isinstance(output, tuple): output = output[0] values = logits_metrics(output[0, -1, :].detach()) raw_logits_rows.append( {key: value.item() for key, value in values.items()} ) output_layer = model.get_output_embeddings() hook = output_layer.register_forward_hook(capture_raw_logits) try: with torch.inference_mode(): result = model.generate( **encoded, do_sample=False, num_beams=1, max_new_tokens=max_new_tokens, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.pad_token_id, repetition_penalty=repetition_penalty, use_cache=True, return_dict_in_generate=True, output_scores=True, ) finally: hook.remove() generated_ids = result.sequences[0, prompt_len:].detach().cpu().tolist() generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True) if len(raw_logits_rows) != len(generated_ids) or len(result.scores) != len(generated_ids): raise RuntimeError( f"step mismatch: ids={len(generated_ids)}, raw={len(raw_logits_rows)}, " f"processed={len(result.scores)}" ) metrics_rows = [] for step, (token_id, raw, processed_scores) in enumerate( zip(generated_ids, raw_logits_rows, result.scores), start=1 ): processed = logits_metrics(processed_scores[0].detach()) metrics_rows.append( { "predicts_token_step": step, "context_generated_tokens": step - 1, "generated_token_id": token_id, "generated_piece": tokenizer.convert_ids_to_tokens(token_id), "raw_top1_probability": raw["top1_probability"], "raw_entropy_nats": raw["entropy_nats"], "raw_top1_top2_margin": raw["top1_top2_margin"], "raw_top1_token_id": raw["top1_token_id"], "processed_top1_probability": processed["top1_probability"].item(), "processed_entropy_nats": processed["entropy_nats"].item(), "processed_top1_top2_margin": processed["top1_top2_margin"].item(), "processed_argmax_token_id": processed["top1_token_id"].item(), "processed_argmax_matches_generated": int( processed["top1_token_id"].item() == token_id ), "rolling_dup8_rate_64": rolling_dup8_rate(generated_ids[:step]), } ) ngram_len, repeat_first, repeat_again = largest_repeated_ngram(generated_ids) eos_emitted = bool(generated_ids and generated_ids[-1] == tokenizer.eos_token_id) summary = { "model": MODEL_ID, "condition": condition, "system_message": messages[0]["content"] if messages[0]["role"] == "system" else None, "repetition_penalty": repetition_penalty, "do_sample": False, "max_new_tokens": max_new_tokens, "eos_token_id": tokenizer.eos_token_id, "prompt_tokens": prompt_len, "generated_tokens": len(generated_ids), "eos_emitted": eos_emitted, "refusal_phrase_present": any( phrase in generated_text for phrase in ("모르겠습니다", "I don't know", "I do not know") ), "largest_repeated_ngram_tokens": ngram_len, "repeat_first_step": repeat_first, "repeat_again_step": repeat_again, "rolling_dup8_mean": sum(row["rolling_dup8_rate_64"] for row in metrics_rows) / max(1, len(metrics_rows)), "rolling_dup8_max": max( (row["rolling_dup8_rate_64"] for row in metrics_rows), default=0.0 ), "rolling_dup8_max_step": max( metrics_rows, key=lambda row: row["rolling_dup8_rate_64"], default={"predicts_token_step": None}, )["predicts_token_step"], "mean_raw_top1_probability": sum(row["raw_top1_probability"] for row in metrics_rows) / max(1, len(metrics_rows)), "final_raw_top1_probability": metrics_rows[-1]["raw_top1_probability"] if metrics_rows else None, "mean_raw_entropy_nats": sum(row["raw_entropy_nats"] for row in metrics_rows) / max(1, len(metrics_rows)), "final_raw_entropy_nats": metrics_rows[-1]["raw_entropy_nats"] if metrics_rows else None, "mean_processed_top1_probability": sum( row["processed_top1_probability"] for row in metrics_rows ) / max(1, len(metrics_rows)), "mean_processed_entropy_nats": sum( row["processed_entropy_nats"] for row in metrics_rows ) / max(1, len(metrics_rows)), } out_dir.mkdir(parents=True, exist_ok=True) (out_dir / "generation.txt").write_text(generated_text, encoding="utf-8") (out_dir / "run_config.json").write_text( json.dumps( { **summary, "messages": messages, "chat_template": tokenizer.chat_template, "generation_uses_model_eos": True, "output_scores": True, "raw_logits_capture": "forward hook on lm_head before generation processors", }, ensure_ascii=False, indent=2, ) + "\n", encoding="utf-8", ) if metrics_rows: with (out_dir / "token_metrics.csv").open("w", encoding="utf-8-sig", newline="") as f: writer = csv.DictWriter(f, fieldnames=list(metrics_rows[0])) writer.writeheader() writer.writerows(metrics_rows) del result torch.cuda.empty_cache() return summary def main(): parser = argparse.ArgumentParser() parser.add_argument( "--model-dir", default="/workspace/rwkv_doomloop_experiment/models/OLMo-2-0425-1B-Instruct", ) parser.add_argument( "--out-dir", default="/workspace/rwkv_doomloop_experiment/runs" ) parser.add_argument("--max-new-tokens", type=int, default=512) args = parser.parse_args() model_dir = Path(args.model_dir) out_root = Path(args.out_dir) tokenizer = AutoTokenizer.from_pretrained(model_dir, local_files_only=True) model = AutoModelForCausalLM.from_pretrained( model_dir, dtype=torch.bfloat16, low_cpu_mem_usage=True, local_files_only=True, ).to("cuda") model.eval() conditions = [ ( "04_olmo2_1b_baseline_greedy", [{"role": "user", "content": USER_PROMPT}], 1.0, ), ( "05_olmo2_1b_system_prompt", [ {"role": "system", "content": SYSTEM_GIVE_UP}, {"role": "user", "content": USER_PROMPT}, ], 1.0, ), ( "06_olmo2_1b_system_plus_repetition_penalty_1_1", [ {"role": "system", "content": SYSTEM_GIVE_UP}, {"role": "user", "content": USER_PROMPT}, ], 1.1, ), ] summaries = [] for name, messages, penalty in conditions: print(f"Starting {name} (penalty={penalty})", flush=True) summary = run_one( model=model, tokenizer=tokenizer, condition=name, messages=messages, repetition_penalty=penalty, max_new_tokens=args.max_new_tokens, out_dir=out_root / name, ) summaries.append(summary) print(json.dumps(summary, ensure_ascii=False), flush=True) (out_root / "olmo2_1b_comparison.json").write_text( json.dumps(summaries, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" ) if __name__ == "__main__": main()