Download scripts/run_olmo2_comparison.py from Ilikemechuri/rwkv_doomloop_experiment: direct link, hf CLI and curl.
- Browser
- Download file 10.1 kB
-
https://huggingface.co/Ilikemechuri/rwkv_doomloop_experiment/resolve/main/scripts/run_olmo2_comparison.py
- Command line
-
hf download hf://Ilikemechuri/rwkv_doomloop_experiment/scripts/run_olmo2_comparison.py
-
curl -L -o run_olmo2_comparison.py https://huggingface.co/Ilikemechuri/rwkv_doomloop_experiment/resolve/main/scripts/run_olmo2_comparison.py
10.1 kB
| #!/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() | |