rwkv_doomloop_experiment / scripts /build_olmo2_artifacts.py
Ilikemechuri's picture
Upload folder using huggingface_hub (part 2)
c2bb3d8 verified
Raw History Blame Contribute Delete
16.9 kB
#!/usr/bin/env python
"""Build OLMo diagnostics and the aligned RWKV/OLMo comparison table."""
import csv
import json
from pathlib import Path
import matplotlib.pyplot as plt
ROOT = Path("/workspace/rwkv_doomloop_experiment")
RUNS = ROOT / "runs"
OLMO_RUNS = [
"04_olmo2_1b_baseline_greedy",
"05_olmo2_1b_system_prompt",
"06_olmo2_1b_system_plus_repetition_penalty_1_1",
"08_olmo2_1b_english_system_prompt",
]
OLMO_CONDITIONS = {
"04_olmo2_1b_baseline_greedy": "baseline_greedy",
"05_olmo2_1b_system_prompt": "system_prompt",
"06_olmo2_1b_system_plus_repetition_penalty_1_1": "system_plus_repetition_penalty_1.1",
"english_system_prompt": "system_prompt_english",
}
def read_csv(path):
with path.open(encoding="utf-8-sig", newline="") as f:
return list(csv.DictReader(f))
def mean(values):
values = [float(value) for value in values if value not in (None, "")]
return sum(values) / len(values) if values else ""
def make_olmo_plots():
for name in OLMO_RUNS:
run_dir = RUNS / name
rows = read_csv(run_dir / "token_metrics.csv")
config = json.loads((run_dir / "run_config.json").read_text(encoding="utf-8"))
x = [int(row["predicts_token_step"]) for row in rows]
fig, axes = plt.subplots(3, 1, figsize=(12, 9), sharex=True)
axes[0].plot(x, [float(r["raw_top1_probability"]) for r in rows], label="raw")
if config["repetition_penalty"] != 1.0:
axes[0].plot(
x,
[float(r["processed_top1_probability"]) for r in rows],
label="after repetition penalty",
alpha=0.8,
)
axes[0].set_ylabel("Top-1 probability")
axes[0].set_ylim(0, 1.02)
axes[0].legend(loc="best")
axes[1].plot(x, [float(r["raw_entropy_nats"]) for r in rows], label="raw")
if config["repetition_penalty"] != 1.0:
axes[1].plot(
x,
[float(r["processed_entropy_nats"]) for r in rows],
label="after repetition penalty",
alpha=0.8,
)
axes[1].set_ylabel("Entropy (nats)")
axes[1].legend(loc="best")
axes[2].plot(
x,
[float(r["rolling_dup8_rate_64"]) for r in rows],
color="tab:red",
)
axes[2].set_ylabel("Repeated 8-gram rate\n(last 64 tokens)")
axes[2].set_xlabel("Generated token step")
axes[2].set_ylim(bottom=0)
first = config.get("repeat_first_step")
again = config.get("repeat_again_step")
length = config.get("largest_repeated_ngram_tokens", 0)
if first and again and length:
for axis in axes:
axis.axvspan(first, first + length - 1, color="orange", alpha=0.12)
axis.axvspan(again, again + length - 1, color="orange", alpha=0.12)
fig.suptitle(f"OLMo 2 1B Instruct — {name}")
fig.tight_layout()
fig.savefig(run_dir / "diagnostics.png", dpi=150, bbox_inches="tight")
plt.close(fig)
def build_comparison():
rwkv_summaries = {
row["condition"]: row
for row in read_csv(ROOT / "comparison.csv")
}
rwkv_run_names = {
"baseline_greedy": "01_baseline_greedy",
"system_prompt": "02_system_prompt",
"system_plus_repetition_penalty_1.1": "03_system_plus_repetition_penalty_1_1",
}
rows_out = []
for condition, run_name in rwkv_run_names.items():
old = rwkv_summaries[condition]
token_rows = read_csv(RUNS / run_name / "token_metrics.csv")
if condition == "system_plus_repetition_penalty_1.1":
raw_top = [row["raw_top1_probability"] for row in token_rows]
raw_ent = [row["raw_entropy_nats"] for row in token_rows]
proc_top = [row["penalty_top1_probability"] for row in token_rows]
proc_ent = [row["penalty_entropy_nats"] for row in token_rows]
else:
raw_top = [row["top1_probability"] for row in token_rows]
raw_ent = [row["entropy_nats"] for row in token_rows]
proc_top, proc_ent = raw_top, raw_ent
rows_out.append(
{
"model": "RWKV/RWKV7-G1j-1.5B-20260831",
"condition": condition,
"generated_tokens": old["generated_tokens"],
"eos_emitted": old["eos_emitted"],
"refusal_phrase_present": old["refusal_phrase_present"],
"largest_repeated_ngram_tokens": old["largest_repeated_ngram_tokens"],
"repeat_first_step": old["repeat_first_step"],
"repeat_again_step": old["repeat_again_step"],
"rolling_dup8_mean": old["rolling_dup8_mean"],
"rolling_dup8_max": old["rolling_dup8_max"],
"rolling_dup8_max_step": old["rolling_dup8_max_step"],
"mean_raw_top1_probability": mean(raw_top),
"final_raw_top1_probability": raw_top[-1],
"mean_raw_entropy_nats": mean(raw_ent),
"final_raw_entropy_nats": raw_ent[-1],
"mean_processed_top1_probability": mean(proc_top),
"mean_processed_entropy_nats": mean(proc_ent),
"mean_wkv_rms": old["mean_wkv_rms"],
"mean_attention_shift_rms": old["mean_attention_shift_rms"],
"mean_ffn_shift_rms": old["mean_ffn_shift_rms"],
"repeated_formula_line_starts": old["repeated_formula_line_starts"],
}
)
olmo_summaries = json.loads((RUNS / "olmo2_1b_comparison.json").read_text(encoding="utf-8"))
for summary in olmo_summaries:
rows_out.append(
{
"model": summary["model"],
"condition": OLMO_CONDITIONS[summary["condition"]],
"generated_tokens": summary["generated_tokens"],
"eos_emitted": summary["eos_emitted"],
"refusal_phrase_present": summary["refusal_phrase_present"],
"largest_repeated_ngram_tokens": summary["largest_repeated_ngram_tokens"],
"repeat_first_step": summary["repeat_first_step"],
"repeat_again_step": summary["repeat_again_step"],
"rolling_dup8_mean": summary["rolling_dup8_mean"],
"rolling_dup8_max": summary["rolling_dup8_max"],
"rolling_dup8_max_step": summary["rolling_dup8_max_step"],
"mean_raw_top1_probability": summary["mean_raw_top1_probability"],
"final_raw_top1_probability": summary["final_raw_top1_probability"],
"mean_raw_entropy_nats": summary["mean_raw_entropy_nats"],
"final_raw_entropy_nats": summary["final_raw_entropy_nats"],
"mean_processed_top1_probability": summary[
"mean_processed_top1_probability"
],
"mean_processed_entropy_nats": summary["mean_processed_entropy_nats"],
"mean_wkv_rms": "",
"mean_attention_shift_rms": "",
"mean_ffn_shift_rms": "",
"repeated_formula_line_starts": "",
}
)
english_summaries = json.loads(
(ROOT / "english_system_comparison.json").read_text(encoding="utf-8")
)
for summary in english_summaries:
rows_out.append(
{
"model": summary["model"],
"condition": "system_prompt_english",
"generated_tokens": summary["generated_tokens"],
"eos_emitted": summary["eos_emitted"],
"refusal_phrase_present": summary["refusal_phrase_present"],
"largest_repeated_ngram_tokens": summary["largest_repeated_ngram_tokens"],
"repeat_first_step": summary["repeat_first_step"],
"repeat_again_step": summary["repeat_again_step"],
"rolling_dup8_mean": summary["rolling_dup8_mean"],
"rolling_dup8_max": summary["rolling_dup8_max"],
"rolling_dup8_max_step": summary.get("rolling_dup8_max_step", ""),
"mean_raw_top1_probability": summary["mean_raw_top1_probability"],
"final_raw_top1_probability": summary["final_raw_top1_probability"],
"mean_raw_entropy_nats": summary["mean_raw_entropy_nats"],
"final_raw_entropy_nats": summary["final_raw_entropy_nats"],
"mean_processed_top1_probability": summary.get(
"mean_processed_top1_probability", summary["mean_raw_top1_probability"]
),
"mean_processed_entropy_nats": summary.get(
"mean_processed_entropy_nats", summary["mean_raw_entropy_nats"]
),
"mean_wkv_rms": summary.get("mean_wkv_rms", ""),
"mean_attention_shift_rms": summary.get("mean_attention_shift_rms", ""),
"mean_ffn_shift_rms": summary.get("mean_ffn_shift_rms", ""),
"repeated_formula_line_starts": "",
}
)
out = ROOT / "comparison_all_models.csv"
with out.open("w", encoding="utf-8-sig", newline="") as f:
writer = csv.DictWriter(f, fieldnames=list(rows_out[0]))
writer.writeheader()
writer.writerows(rows_out)
build_language_comparison(english_summaries)
def build_language_comparison(english_summaries):
rwkv_korean_summary = next(
row
for row in read_csv(ROOT / "comparison.csv")
if row["condition"] == "system_prompt"
)
by_model_and_language = []
configs = [
(
"RWKV/RWKV7-G1j-1.5B-20260831",
"Korean",
RUNS / "02_system_prompt",
None,
rwkv_korean_summary,
),
(
"RWKV/RWKV7-G1j-1.5B-20260831",
"English",
RUNS / "07_rwkv_english_system_prompt",
english_summaries[0],
None,
),
(
"allenai/OLMo-2-0425-1B-Instruct",
"Korean",
RUNS / "05_olmo2_1b_system_prompt",
None,
None,
),
(
"allenai/OLMo-2-0425-1B-Instruct",
"English",
RUNS / "08_olmo2_1b_english_system_prompt",
english_summaries[1],
None,
),
]
for model_id, language, run_dir, english_summary, fallback_summary in configs:
config = json.loads((run_dir / "run_config.json").read_text(encoding="utf-8"))
summary = english_summary or fallback_summary or config
by_model_and_language.append(
{
"model": model_id,
"system_message_language": language,
"prompt_tokens": config.get(
"prompt_tokens", 122 if model_id.startswith("RWKV/") else ""
),
"generated_tokens": summary["generated_tokens"],
"eos_emitted": summary["eos_emitted"],
"refusal_phrase_present": summary["refusal_phrase_present"],
"largest_repeated_ngram_tokens": summary[
"largest_repeated_ngram_tokens"
],
"repeat_first_step": summary["repeat_first_step"],
"repeat_again_step": summary["repeat_again_step"],
"rolling_dup8_mean": summary["rolling_dup8_mean"],
"rolling_dup8_max": summary["rolling_dup8_max"],
"mean_raw_top1_probability": summary.get(
"mean_raw_top1_probability", summary.get("mean_top1_probability", "")
),
"final_raw_top1_probability": summary.get(
"final_raw_top1_probability", summary.get("final_top1_probability", "")
),
"mean_raw_entropy_nats": summary.get(
"mean_raw_entropy_nats", summary.get("mean_entropy_nats", "")
),
"final_raw_entropy_nats": summary.get(
"final_raw_entropy_nats", summary.get("final_entropy_nats", "")
),
"mean_wkv_rms": summary.get("mean_wkv_rms", ""),
"final_wkv_rms": summary.get("final_wkv_rms", ""),
"mean_attention_shift_rms": summary.get("mean_attention_shift_rms", ""),
"mean_ffn_shift_rms": summary.get("mean_ffn_shift_rms", ""),
}
)
with (ROOT / "english_system_language_comparison.csv").open(
"w", encoding="utf-8-sig", newline=""
) as f:
writer = csv.DictWriter(f, fieldnames=list(by_model_and_language[0]))
writer.writeheader()
writer.writerows(by_model_and_language)
def make_rwkv_english_plot():
run_dir = RUNS / "07_rwkv_english_system_prompt"
rows = read_csv(run_dir / "token_metrics.csv")
config = json.loads((run_dir / "run_config.json").read_text(encoding="utf-8"))
x = [int(row["predicts_token_step"]) for row in rows]
fig, axes = plt.subplots(4, 1, figsize=(12, 11), sharex=True)
axes[0].plot(x, [float(r["raw_top1_probability"]) for r in rows], label="raw")
axes[0].set_ylabel("Top-1 probability")
axes[0].set_ylim(0, 1.02)
axes[1].plot(x, [float(r["raw_entropy_nats"]) for r in rows], color="tab:orange")
axes[1].set_ylabel("Entropy (nats)")
axes[2].plot(
x,
[float(r["rolling_dup8_rate_64"]) for r in rows],
color="tab:red",
)
axes[2].set_ylabel("Repeated 8-gram rate\n(last 64 tokens)")
axes[2].set_ylim(bottom=0)
axes[3].plot(x, [float(r["wkv_rms_mean"]) for r in rows], label="WKV")
axes[3].plot(x, [float(r["att_shift_rms_mean"]) for r in rows], label="attention shift")
axes[3].plot(x, [float(r["ffn_shift_rms_mean"]) for r in rows], label="FFN shift")
axes[3].set_ylabel("State RMS")
axes[3].set_xlabel("Generated token step")
axes[3].legend(loc="best")
first = config.get("repeat_first_step")
again = config.get("repeat_again_step")
length = config.get("largest_repeated_ngram_tokens", 0)
if first and again and length:
for axis in axes:
axis.axvspan(first, first + length - 1, color="orange", alpha=0.12)
axis.axvspan(again, again + length - 1, color="orange", alpha=0.12)
fig.suptitle("RWKV-7 — English system message")
fig.tight_layout()
fig.savefig(run_dir / "diagnostics.png", dpi=150, bbox_inches="tight")
plt.close(fig)
def update_manifest():
manifest_path = ROOT / "run_manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
model_dir = ROOT / "models" / "OLMo-2-0425-1B-Instruct"
manifest["comparison_model"] = {
"model": "allenai/OLMo-2-0425-1B-Instruct",
"model_revision": "48d788eca847d4d7548f375ad03d3c9312f6139e",
"model_card": "https://huggingface.co/allenai/OLMo-2-0425-1B-Instruct",
"local_snapshot": str(model_dir),
"model_safetensors_bytes": (model_dir / "model.safetensors").stat().st_size,
"dtype": "bfloat16",
"eos_token_id": 100257,
"pad_token_id": 100277,
"chat_template": "tokenizer.apply_chat_template; BOS + system/user roles + assistant generation prompt",
"generation_conditions": [
"greedy baseline, no system prompt",
"greedy with same give-up system prompt",
"greedy with same give-up system prompt and repetition_penalty=1.1",
],
"instrumentation": {
"raw_logits": "lm_head forward hook before generation processors",
"processed_scores": "generate(output_scores=True), after generation processors",
"repetition": "rolling duplicate 8-gram rate in last up to 64 model tokens; largest exact non-overlapping repeated token span",
"state_norm": "not applicable; OLMo 2 is a Transformer and has no RWKV recurrent state",
},
"workspace_is_volume": False,
}
manifest["english_system_comparison"] = {
"english_system_message": "If you cannot solve the problem or do not know the answer, do not keep guessing. Say “I don't know” and stop.",
"user_prompt_language": "Korean, same mathematical problem as previous runs",
"do_sample": False,
"max_new_tokens": 512,
"repetition_penalty": 1.0,
"runs": [
"runs/07_rwkv_english_system_prompt",
"runs/08_olmo2_1b_english_system_prompt",
],
"note": "One deterministic generation per model; system-message language is the only intended prompt change relative to the earlier Korean system-message condition.",
}
manifest_path.write_text(
json.dumps(manifest, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
if __name__ == "__main__":
make_olmo_plots()
make_rwkv_english_plot()
build_comparison()
update_manifest()