diff options
Diffstat (limited to 'worldalign/evaluate_energy_prefix.py')
| -rw-r--r-- | worldalign/evaluate_energy_prefix.py | 151 |
1 files changed, 151 insertions, 0 deletions
diff --git a/worldalign/evaluate_energy_prefix.py b/worldalign/evaluate_energy_prefix.py new file mode 100644 index 0000000..2f3af7a --- /dev/null +++ b/worldalign/evaluate_energy_prefix.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from .common import dtype_for_device, write_json +from .evaluate import unigram_f1 +from .io import load_feature_pair, select_rows +from .models import load_prefix + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--energy", default="artifacts/energy_free_test.pt") + parser.add_argument("--vision", default="artifacts/vision.pt") + parser.add_argument("--text", default="artifacts/text.pt") + parser.add_argument("--text-orbits") + parser.add_argument("--prefix", default="artifacts/prefix.pt") + parser.add_argument("--samples", type=int, default=100) + parser.add_argument("--max-new-tokens", type=int, default=32) + parser.add_argument("--device", default="cuda:3") + parser.add_argument( + "--output", default="artifacts/energy_prefix_evaluation.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + energy = torch.load(args.energy, map_location="cpu", weights_only=False) + _, text, _, text_lookup = load_feature_pair(args.vision, args.text) + rows = energy["rows"][: args.samples] + initial = energy["initial_particles"][: args.samples] + final = energy["final_particles"][: args.samples] + oracle = select_rows(text["features"], text_lookup, rows) + if args.text_orbits: + orbit_state = torch.load( + args.text_orbits, map_location="cpu", weights_only=False + ) + orbit_lookup = { + int(row): index + for index, row in enumerate(orbit_state["rows"]) + } + orbit_mean = torch.nn.functional.normalize( + orbit_state["features"].float().mean(1), dim=-1 + ) + oracle = select_rows(orbit_mean, orbit_lookup, rows) + reference_lookup = { + int(row): captions + for row, captions in zip(text["rows"], text["all_captions"]) + } + + prefix, prefix_state = load_prefix(args.prefix, args.device) + if prefix.semantic_dim != initial.shape[-1]: + raise ValueError("Energy latent and text-only prefix dimensions differ") + tokenizer = AutoTokenizer.from_pretrained(text["model"]) + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + dtype = dtype_for_device(args.device) + lm = AutoModelForCausalLM.from_pretrained( + text["model"], torch_dtype=dtype + ).to(args.device) + lm.eval() + + conditions = {"initial": initial, "final": final, "oracle": oracle} + decoded: dict[str, list[str]] = {key: [] for key in conditions} + with torch.inference_mode(): + for name, semantic in conditions.items(): + for chunk in semantic.split(16): + embeds = prefix(chunk.to(args.device)).to(dtype) + attention = torch.ones( + embeds.shape[:2], dtype=torch.long, device=args.device + ) + output = lm.generate( + inputs_embeds=embeds, + attention_mask=attention, + max_new_tokens=args.max_new_tokens, + do_sample=False, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + ) + decoded[name].extend( + tokenizer.batch_decode( + output, skip_special_tokens=True + ) + ) + + scores = {} + per_condition: dict[str, list[float]] = {} + for name, generations in decoded.items(): + values = [ + unigram_f1(generation, reference_lookup[int(row)]) + for row, generation in zip(rows, generations) + ] + per_condition[name] = values + scores[name] = sum(values) / max(len(values), 1) + difference = np.asarray(per_condition["final"]) - np.asarray( + per_condition["initial"] + ) + bootstrap_generator = np.random.default_rng(20260729) + bootstrap = difference[ + bootstrap_generator.integers( + 0, len(difference), size=(10_000, len(difference)) + ) + ].mean(1) + result = { + "samples": len(rows), + "mean_best_reference_unigram_f1": scores, + "paired_final_minus_initial": { + "mean": float(difference.mean()), + "bootstrap_95_percentile_interval": [ + float(np.quantile(bootstrap, 0.025)), + float(np.quantile(bootstrap, 0.975)), + ], + "improved": int((difference > 0).sum()), + "tied": int((difference == 0).sum()), + "worsened": int((difference < 0).sum()), + }, + "energy_protocol": energy["protocol"], + "prefix_training": prefix_state["training"], + "per_sample_unigram_f1": [ + { + "row": int(row), + **{ + name: per_condition[name][index] + for name in per_condition + }, + } + for index, row in enumerate(rows) + ], + "examples": [ + { + "row": int(row), + "references": reference_lookup[int(row)], + **{name: decoded[name][i] for name in decoded}, + } + for i, row in enumerate(rows[:25]) + ], + } + Path(args.output).parent.mkdir(parents=True, exist_ok=True) + write_json(args.output, result) + print(json.dumps(result, indent=2, ensure_ascii=False)) + + +if __name__ == "__main__": + main() |
