rwkv_doomloop_experiment / scripts /run_olmo2_comparison.py
Ilikemechuri's picture
Upload folder using huggingface_hub (part 2)
c2bb3d8 verified
Raw History Blame Contribute Delete
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()