""" Unified hallucination-removal validation script. Supports any scene→object relation defined in experiment/config/relations.json. Use --relation to select (default: bathroom_toilet). Supported model types (--model_type): lora -> loads base model + LoRA adapter via PeftModel merged -> loads merged HF model via AutoModelForPreTraining delta_w -> loads HookedSAELlavaConditionalGeneration + .pt state dict grace -> loads base model + restores GRACE codebook adapters wise -> loads base model + restores WISE adapter state dualedit -> loads base model + restores DualEdit adapters visedit -> loads editor with trained checkpoint Mention detection (--mention_method): keyword -> fast negation-aware regex (same as old scripts) llm -> local LLM judge only both -> keyword + LLM side-by-side Usage: # LoRA adapter with custom relation python -m experiment.evaluation.validate \ --relation kitchen_microwave \ --model_type lora \ --model_dir step3_lora_v5_outputs/kitchen_microwave/run_xxx/lora_adapter # Default (bathroom_toilet) for backward compat python -m experiment.evaluation.validate \ --model_type lora \ --model_dir step3_lora_outputs/lora_adapter """ from __future__ import annotations import os import sys import csv import json import math import argparse from pathlib import Path # Must be set before importing vllm — subprocess inherits env at spawn time os.environ["VLLM_USE_V1"] = "0" os.environ.setdefault("NCCL_P2P_DISABLE", "1") os.environ.setdefault("NCCL_IB_DISABLE", "1") import torch sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../..")) from transformers import AutoProcessor, AutoModelForPreTraining from experiment.config.relation_config import get_relation_config, RelationConfig from experiment.data.datasets import get_split_image_ids from experiment.data.hf_loader import load_hf_dataset from experiment.evaluation.metrics import build_metrics from experiment.evaluation.inference import ( collect_outputs_transformers, collect_outputs_vllm, collect_outputs_visedit, ) from experiment.evaluation.metric import evaluate_collected_outputs, compute_kme_metrics from experiment.evaluation.metrics import TextSimilarityScorer from experiment.evaluation.summary import print_summary MODEL_NAME = "llava-hf/llava-1.5-7b-hf" dtype = torch.float16 def parse_args(): parser = argparse.ArgumentParser(description="Unified hallucination-removal validation") parser.add_argument("--relation", type=str, default="bathroom_toilet", help="Relation key from relations.json (default: bathroom_toilet)") parser.add_argument("--val_csv", type=str, default=None, help="(Legacy) Path to CSV. If omitted, loads from HuggingFace.") parser.add_argument("--val_image_dir", type=str, default=None, help="(Legacy) Image directory. If omitted, loads from HuggingFace.") parser.add_argument("--dataset_id", type=str, default=None, help="HuggingFace dataset ID (default: auto from relation config)") parser.add_argument("--num_per_category", type=int, default=50) parser.add_argument("--use_val_split", action="store_true", help="Only use val-split images") parser.add_argument("--prompts", type=str, nargs="+", default=["Describe this image.", "What do you see in this image?"]) parser.add_argument("--generality_prompts", type=str, nargs="*", default=["Give a detailed description of this image."], help="Unseen prompts for generality evaluation (not used in training)") parser.add_argument("--train_prompts", type=str, nargs="*", default=None, help="Prompts that were used during training (default: from relation config)") parser.add_argument("--max_new_tokens", type=int, default=300) parser.add_argument("--model_type", type=str, choices=["lora", "merged", "delta_w", "grace", "wise", "dualedit", "visedit"], required=True, help="How to load the finetuned model") parser.add_argument("--base_model_name", type=str, default=MODEL_NAME, help="HuggingFace model name for the base / original model") group = parser.add_mutually_exclusive_group(required=True) group.add_argument("--model_dir", type=str, help="Path to merged model or LoRA adapter dir (lora / merged)") group.add_argument("--checkpoint", type=str, help="Path to .pt state dict (delta_w)") parser.add_argument("--mention_method", type=str, choices=["keyword", "llm", "both"], default="both") parser.add_argument("--clip_model", type=str, default="google/siglip-base-patch16-224") parser.add_argument("--judge_model", type=str, default="Qwen/Qwen3-VL-32B-Instruct") parser.add_argument("--judge_device", type=str, default="cuda", help="Device for judge, e.g. cuda, cuda:1, or cpu") parser.add_argument("--judge_max_tokens", type=int, default=150) parser.add_argument("--inference_backend", type=str, choices=["transformers", "vllm"], default="transformers") parser.add_argument("--vllm_batch_size", type=int, default=64) parser.add_argument("--vllm_tensor_parallel_size", type=int, default=1) parser.add_argument("--vllm_gpu_memory_utilization", type=float, default=0.9) parser.add_argument("--vllm_max_model_len", type=int, default=4096) parser.add_argument("--output_dir", type=str, default="./step4_v2_outputs") parser.add_argument("--original_cache_dir", type=str, default="./cached_original_outputs", help="Directory to cache original model outputs for reuse") parser.add_argument("--edit_image_ids", type=str, default=None, help="(deprecated, ignored) Previously pinned BNT eval to specific IDs.") parser.add_argument("--edit_targets", type=str, default=None, help="(visedit only) Path to eval_targets.json {image_id: target_new} " "written by run_visedit.py. Used as correction target for edit signal.") parser.add_argument("--visedit_dir", type=str, default=None, help="(visedit only) Path to VisEdit repo root. " "Defaults to /VisEdit.") return parser.parse_args() def load_category_images(relation_config: RelationConfig, num_per_category, use_val_split=False, dataset_id=None, csv_path=None, image_dir=None, edit_image_ids=None): """Load first num_per_category images per category deterministically.""" dataset_id = dataset_id or relation_config.dataset_id scene_col = relation_config.scene_key object_col = relation_config.object_key categories = {name: [] for name in relation_config.category_names} if csv_path is not None and image_dir is not None: # Legacy CSV loading val_ids = get_split_image_ids(csv_path, "val") if use_val_split else None with open(csv_path, "r") as f: reader = csv.DictReader(f) for row in reader: if val_ids is not None and row["image_id"] not in val_ids: continue image_path = os.path.join(image_dir, f"{row['image_id']}.jpg") if not os.path.exists(image_path): continue entry = {"path": image_path, "image_id": row["image_id"]} scene_val = int(row.get(scene_col, 0)) object_val = int(row.get(object_col, 0)) cat = _classify(scene_val, object_val, relation_config) if cat in categories and len(categories[cat]) < num_per_category: categories[cat].append(entry) if all(len(v) >= num_per_category for v in categories.values()): break else: # HuggingFace dataset if use_val_split: ds = load_hf_dataset(dataset_id, split="val") else: ds = load_hf_dataset(dataset_id) if hasattr(ds, "keys"): from datasets import concatenate_datasets ds = concatenate_datasets([ds[s] for s in ds]) for item in ds: scene_val = int(item[scene_col]) object_val = int(item[object_col]) entry = {"image": item["image"], "image_id": item["image_id"]} cat = _classify(scene_val, object_val, relation_config) if cat in categories and len(categories[cat]) < num_per_category: categories[cat].append(entry) if all(len(v) >= num_per_category for v in categories.values()): break for cat, imgs in categories.items(): print(f" {cat}: {len(imgs)} images") return categories def _classify(scene_val: int, object_val: int, rc: RelationConfig) -> str: """Classify an image into one of the 4 categories.""" if scene_val == 1 and object_val == 0: return rc.scene_no_object elif scene_val == 1 and object_val == 1: return rc.scene_with_object elif scene_val == 0 and object_val == 1: return rc.non_scene_with_object else: return "unrelated" def _is_adapter_dir(path: str) -> bool: # Local directory with adapter_config.json if os.path.exists(os.path.join(path, "adapter_config.json")): return True # HuggingFace Hub repo ID (e.g. "user/repo-name") if not os.path.isabs(path) and path.count("/") == 1: return True return False def _restore_adapters(model, adapter_states, device): """Reconstruct GRACE/WISE adapter modules from saved state.""" import copy for module_path, saved in adapter_states.items(): parent_path, attr_name = module_path.rsplit(".", 1) parent = model.get_submodule(parent_path) original_layer = getattr(parent, attr_name) extra = saved["extra"] cfg = extra["config"] if saved["type"] == "GRACEAdapter": from easyeditor.models.grace.GRACE import GRACEAdapter config = type("Cfg", (), { "eps": cfg["eps"], "dist_fn": cfg["dist_fn"], "replacement": cfg["replacement"], "num_pert": cfg["num_pert"], "dropout": 0.0, "val_init": cfg.get("val_init", "cold"), })() adapter = GRACEAdapter(config, original_layer, transpose=True).to(device) adapter.keys = extra["keys"].to(device) adapter.values = torch.nn.Parameter( saved["state_dict"]["values"].to(device)) adapter.epsilons = extra["epsilons"].to(device) adapter.key_labels = extra["key_labels"] adapter.edit_ids = extra["edit_ids"] elif saved["type"] == "WISEAdapter": from easyeditor.models.wise.WISE import WISEAdapter config = type("Cfg", (), { "model_name": cfg["model_name"], "retrieve": cfg["retrieve"], "act_ratio": cfg["act_ratio"], "merge_alg": cfg["merge_alg"], "save_freq": cfg.get("save_freq"), "densities": cfg.get("densities"), "weights": cfg.get("weights"), })() adapter = WISEAdapter(config, original_layer, transpose=True).to(device) adapter.new_weight = extra["new_weight"].to(device) adapter.original_layer.load_state_dict(extra["original_layer_state"]) adapter.original_layer = adapter.original_layer.to(device) adapter.memory_weight = [w.to(device) for w in extra["memory_weight"]] adapter.memory_mean_act = extra["memory_mean_act"] adapter.editing_mean_act = extra["editing_mean_act"] # Restore learned parameters from state_dict adapter.load_state_dict(saved["state_dict"], strict=False) else: raise ValueError(f"Unknown adapter type: {saved['type']}") setattr(parent, attr_name, adapter) return model def _restore_dualedit(model, state_path: str, device: str): """Reconstruct DualEdit adapters from saved state and hook them to model.""" from experiment.knowledge_editing.dualedit.adapter import VisionEditAdapter, TextEditAdapter state = torch.load(state_path, map_location=device, weights_only=False) hp = state["hparams"] # Create adapters vision_adapter = VisionEditAdapter( hidden_size=hp["hidden_size"], mid_dim=hp["adapter_mid_dim"], cross_att_head_n=hp["cross_att_head_n"], img_tok_n=hp["img_tok_n"], ).to(device) text_adapter = TextEditAdapter( hidden_size=hp["hidden_size"], mid_dim=hp["adapter_mid_dim"], cross_att_head_n=hp["cross_att_head_n"], ).to(device) # Load trained weights vision_adapter.load_state_dict(state["vision_adapter_state"]) text_adapter.load_state_dict(state["text_adapter_state"]) # Set edit signals (mean over training set) vision_adapter.set_edit_signal( state["mean_vis_edit_reps"].to(device), state["mean_vis_edit_mask"].to(device), ) text_adapter.set_edit_signal( state["mean_txt_edit_reps"].to(device), state["mean_txt_edit_mask"].to(device), ) # Set gate vision_adapter.set_gate(state["gate_prototype"].to(device), state["gate_threshold"]) text_adapter.set_gate(state["gate_prototype"].to(device), state["gate_threshold"]) vision_adapter.open_adapter(True) text_adapter.open_adapter(True) vision_adapter.open_gating = True text_adapter.open_gating = True # Hook adapters to model layers vision_layer_name = hp["llm_layer_tmp"].format(hp["vision_adapter_layer"]) text_layer_name = hp["llm_layer_tmp"].format(hp["text_adapter_layer"]) def _find_module(m, path): for part in path.split("."): m = m[int(part)] if part.isdigit() else getattr(m, part) return m def make_hook(adapter): def hook(module, args, output): if isinstance(output, tuple): out = list(output) out[0] = adapter(out[0]) return tuple(out) return adapter(output) return hook vision_layer = _find_module(model, vision_layer_name) text_layer = _find_module(model, text_layer_name) vision_layer.register_forward_hook(make_hook(vision_adapter)) text_layer.register_forward_hook(make_hook(text_adapter)) # Patch model.generate to call set_input_info before each generation. image_token_id = model.config.image_token_index img_tok_n = hp["img_tok_n"] _va = vision_adapter _ta = text_adapter _original_generate = model.generate def _generate_with_adapter_info(*args, **kwargs): input_ids = kwargs.get("input_ids") if input_ids is not None and input_ids.shape[1] > 1: positions = (input_ids[0] == image_token_id).nonzero(as_tuple=True)[0] if len(positions) > 0: vt_begin = int(positions[0]) vt_end = vt_begin + img_tok_n merged_len = input_ids.shape[1] - 1 + img_tok_n print(f" [DualEdit] set_input_info: vt_begin={vt_begin}, vt_end={vt_end}, merged_len={merged_len}, image_token_id={image_token_id}") _va.set_input_info(True, vt_begin, vt_end) _ta.set_input_info(True, vt_begin, vt_end) _ta.prompt_end = torch.tensor([merged_len], device=input_ids.device) else: print(f" [DualEdit] WARNING: image token {image_token_id} not found in input_ids (tokens: {input_ids[0].tolist()[:10]}...)") _va.set_input_info(False, None, None) _ta.set_input_info(False, None, None) else: print(f" [DualEdit] WARNING: input_ids missing or single-token in generate kwargs") return _original_generate(*args, **kwargs) model.generate = _generate_with_adapter_info # Store refs for potential later access model._dualedit_vision_adapter = vision_adapter model._dualedit_text_adapter = text_adapter return model def load_base_model(base_model_name: str, device: str): model = AutoModelForPreTraining.from_pretrained( base_model_name, torch_dtype=dtype, ).to(device) model.eval() return model def load_finetuned_model(args, device: str): if args.model_type == "lora": from peft import PeftModel model_dir = args.model_dir if _is_adapter_dir(model_dir): print(f" Detected LoRA adapter at {model_dir}") base = AutoModelForPreTraining.from_pretrained( args.base_model_name, torch_dtype=dtype, ).to(device) model = PeftModel.from_pretrained(base, model_dir) else: print(" No adapter_config.json found; treating as merged model") model = AutoModelForPreTraining.from_pretrained( model_dir, torch_dtype=dtype, ).to(device) model.eval() return model if args.model_type == "merged": model = AutoModelForPreTraining.from_pretrained( args.model_dir, torch_dtype=dtype, ).to(device) model.eval() return model if args.model_type == "delta_w": from model.llava.hooked_llava import HookedSAELlavaConditionalGeneration model = HookedSAELlavaConditionalGeneration.from_pretrained( args.base_model_name, torch_dtype=dtype, ).to(device) state_dict = torch.load(args.checkpoint, map_location=device) model.load_state_dict(state_dict, strict=True) model.eval() return model if args.model_type in ("grace", "wise"): model = AutoModelForPreTraining.from_pretrained( args.base_model_name, torch_dtype=dtype, ).to(device) adapter_path = os.path.join(args.model_dir, "adapter_state.pt") if os.path.exists(adapter_path): adapter_states = torch.load(adapter_path, map_location=device, weights_only=False) model = _restore_adapters(model, adapter_states, device) print(f" Restored {len(adapter_states)} adapter(s) from {adapter_path}") model.eval() return model if args.model_type == "dualedit": model = AutoModelForPreTraining.from_pretrained( args.base_model_name, torch_dtype=dtype, ).to(device) dualedit_path = os.path.join(args.model_dir, "dualedit_state.pt") if os.path.exists(dualedit_path): model = _restore_dualedit(model, dualedit_path, device) print(f" Restored DualEdit adapters from {dualedit_path}") model.eval() return model if args.model_type == "visedit": visedit_dir = args.visedit_dir or str( Path(__file__).resolve().parents[2] / "VisEdit" ) if visedit_dir not in sys.path: sys.path.insert(0, visedit_dir) # Patch GLOBAL.py to point at correct root + model path global_py = Path(visedit_dir) / "utils" / "GLOBAL.py" global_py.write_text( f"ROOT_PATH = {visedit_dir!r}\n" f"model_path_map = {{\n" f" 'llava-v1.5-7b': {args.base_model_name!r},\n" f" 'blip2-opt-2.7b': 'models/blip2-opt-2.7b',\n" f" 'minigpt-4-vicuna-7b': 'models/minigpt-4-vicuna-7b',\n" f"}}\n" ) from utils import load_vllm_editor ckpt_path = args.model_dir editor = load_vllm_editor( "vead", "llava", device, extra_devices=[], editor_ckpt_path=ckpt_path, for_train=False, ) print(f" Loaded VEAD editor from {ckpt_path}") return editor raise ValueError(f"Unknown model_type: {args.model_type}") def _load_processor_from(name_or_path: str): """Load processor, falling back to manual component construction on version mismatches.""" try: return AutoProcessor.from_pretrained(name_or_path) except Exception: pass try: from transformers import AutoTokenizer, CLIPImageProcessor, LlavaProcessor tokenizer = AutoTokenizer.from_pretrained(name_or_path, use_fast=False) image_processor = CLIPImageProcessor.from_pretrained(name_or_path) return LlavaProcessor(tokenizer=tokenizer, image_processor=image_processor) except Exception as e: raise RuntimeError( f"Failed to load processor from {name_or_path!r}. " "Try deleting the HuggingFace cache for this model and re-downloading." ) from e def load_processor(args): if args.model_type == "lora" and args.model_dir and _is_adapter_dir(args.model_dir): return _load_processor_from(args.base_model_name) if args.model_type in ("delta_w", "grace", "wise", "dualedit", "visedit"): return _load_processor_from(args.base_model_name) try: source = args.model_dir if args.model_dir else args.base_model_name return _load_processor_from(source) except Exception: return _load_processor_from(args.base_model_name) def infer_run_name(args) -> str: path = args.model_dir or args.checkpoint or "unknown" parts = os.path.normpath(path).split(os.sep) for part in reversed(parts): if part.startswith("run_"): return part return os.path.basename(os.path.dirname(path)) or os.path.basename(path) or "run" def main(): args = parse_args() # Load relation config relation_config = get_relation_config(args.relation) dataset_id = args.dataset_id or relation_config.dataset_id # Resolve train_prompts from relation config if not specified train_prompts = args.train_prompts or relation_config.train_prompts if args.inference_backend == "vllm": device = "cuda" else: device = "cuda" if torch.cuda.is_available() else "cpu" run_name = infer_run_name(args) eval_dir = os.path.join(args.output_dir, run_name) os.makedirs(eval_dir, exist_ok=True) # Merge generality (unseen) prompts into the full prompt list for inference. all_prompts = list(args.prompts) if args.generality_prompts: for p in args.generality_prompts: if p not in all_prompts: all_prompts.append(p) args.prompts = all_prompts print("=" * 70) print("Validate Hallucination Removal") print("=" * 70) print(f" relation: {relation_config}") print(f" model_type: {args.model_type}") print(f" mention_method: {args.mention_method}") print(f" backend: {args.inference_backend}") print(f" device: {device}") print(f" output_dir: {eval_dir}") print(f" train_prompts: {train_prompts}") print(f" all_prompts: {all_prompts}") if args.edit_image_ids: print(f" NOTE: --edit_image_ids is deprecated and ignored (using first N deterministically)") print("\nLoading images by category...") categories = load_category_images( relation_config=relation_config, num_per_category=args.num_per_category, use_val_split=args.use_val_split, dataset_id=dataset_id, csv_path=args.val_csv, image_dir=args.val_image_dir, ) processor = load_processor(args) # -- Original model outputs: load from cache or generate & save -- cache_dir = args.original_cache_dir os.makedirs(cache_dir, exist_ok=True) # Include relation in cache filename to avoid cross-relation collisions all_cache_file = os.path.join(cache_dir, f"original_outputs_{args.relation}_all.json") specific_cache_file = os.path.join( cache_dir, f"original_outputs_{args.relation}_n{args.num_per_category}_p{len(args.prompts)}.json", ) if os.path.exists(all_cache_file): print(f"\n[1/3] Loading cached original model outputs from {all_cache_file}") with open(all_cache_file, "r") as f: all_cache = json.load(f) needed_ids = { cat: {entry["image_id"] for entry in entries} for cat, entries in categories.items() } needed_prompts = set(args.prompts) original_outputs = {} for cat, entries in all_cache.items(): original_outputs[cat] = [ e for e in entries if e["image_id"] in needed_ids.get(cat, set()) and e["prompt"] in needed_prompts ] elif os.path.exists(specific_cache_file): print(f"\n[1/3] Loading cached original model outputs from {specific_cache_file}") with open(specific_cache_file, "r") as f: original_outputs = json.load(f) else: print("\n[1/3] Inference: loading original model...") original_model = load_base_model(args.base_model_name, device) shared_collect_orig = dict( processor=processor, categories=categories, prompts=args.prompts, max_new_tokens=args.max_new_tokens, device=device, ) print(" Collecting original model outputs...") original_outputs = collect_outputs_transformers( model=original_model, label="original", **shared_collect_orig ) del original_model torch.cuda.empty_cache() with open(specific_cache_file, "w") as f: json.dump(original_outputs, f, indent=2) print(f" Cached original outputs to {specific_cache_file}") # -- Fine-tuned model outputs -- shared_collect = dict( processor=processor, categories=categories, prompts=args.prompts, max_new_tokens=args.max_new_tokens, device=device, ) print("\n[1/3] Inference: loading fine-tuned model...") finetuned_model = load_finetuned_model(args, device) print(" Collecting fine-tuned model outputs...") if args.model_type == "visedit": edit_targets = None if args.edit_targets and os.path.exists(args.edit_targets): with open(args.edit_targets) as f: edit_targets = json.load(f) print(f" Loaded {len(edit_targets)} edit targets from {args.edit_targets}") finetuned_outputs = collect_outputs_visedit( editor=finetuned_model, categories=categories, prompts=args.prompts, max_new_tokens=args.max_new_tokens, edit_targets=edit_targets, label="visedit", relation_config=relation_config, ) else: finetuned_outputs = collect_outputs_transformers( model=finetuned_model, label="finetuned", **shared_collect ) del finetuned_model torch.cuda.empty_cache() print("\nLoading metrics...") keyword_detector, clip_scorer, judge = build_metrics( mention_method=args.mention_method, clip_model=args.clip_model, judge_model=args.judge_model, judge_device=args.judge_device, judge_max_tokens=args.judge_max_tokens, mention_keywords=relation_config.mention_keywords, object_name=relation_config.judge_object_name, ) print("\n[2/3] Evaluating metrics from collected outputs...") image_lookup = { entry["image_id"]: entry["image"] for cat_entries in categories.values() for entry in cat_entries if "image" in entry } original_results = evaluate_collected_outputs( collected_outputs=original_outputs, keyword_detector=keyword_detector, clip_scorer=clip_scorer, judge=judge, mention_method=args.mention_method, label="original", image_lookup=image_lookup, ) finetuned_results = evaluate_collected_outputs( collected_outputs=finetuned_outputs, keyword_detector=keyword_detector, clip_scorer=clip_scorer, judge=judge, mention_method=args.mention_method, label="finetuned", image_lookup=image_lookup, ) print("\n Computing KME metrics (locality, generality, consistency)...") similarity_scorer = TextSimilarityScorer() kme_metrics = compute_kme_metrics( original_outputs=original_outputs, finetuned_outputs=finetuned_outputs, keyword_detector=keyword_detector, train_prompts=train_prompts, similarity_scorer=similarity_scorer, efficacy_category=relation_config.efficacy_category, locality_categories=relation_config.locality_categories, ) print("\n[3/3] Results") print_summary(categories, original_results, finetuned_results, kme_metrics=kme_metrics, relation_config=relation_config) summary = {} for cat in categories: o = original_results[cat] f = finetuned_results[cat] summary[cat] = { "original": {k: v for k, v in o.items() if k != "details"}, "finetuned": {k: v for k, v in f.items() if k != "details"}, "delta_clip": ( (f["avg_clip_score"] - o["avg_clip_score"]) if not (math.isnan(f["avg_clip_score"]) or math.isnan(o["avg_clip_score"])) else None ), } config_snapshot = { "relation": args.relation, "model_type": args.model_type, "base_model_name": args.base_model_name, "model_dir": args.model_dir, "checkpoint": args.checkpoint, "val_csv": args.val_csv, "num_per_category": args.num_per_category, "prompts": args.prompts, "train_prompts": train_prompts, "generality_prompts": args.generality_prompts, "mention_method": args.mention_method, "clip_model": args.clip_model, "judge_model": args.judge_model, "judge_device": args.judge_device, "inference_backend": args.inference_backend, "vllm_batch_size": args.vllm_batch_size, "vllm_tensor_parallel_size": args.vllm_tensor_parallel_size, "vllm_gpu_memory_utilization": args.vllm_gpu_memory_utilization, "vllm_max_model_len": args.vllm_max_model_len, } # Serialize KME metrics (NaN → null for JSON) kme_serializable = { k: (None if isinstance(v, float) and math.isnan(v) else v) for k, v in kme_metrics.items() } results_path = os.path.join(eval_dir, "validation_results.json") with open(results_path, "w") as f: json.dump({ "summary": summary, "kme_metrics": kme_serializable, "config": config_snapshot, }, f, indent=2) details_path = os.path.join(eval_dir, "validation_details.json") with open(details_path, "w") as f: json.dump({ "original": {cat: r["details"] for cat, r in original_results.items()}, "finetuned": {cat: r["details"] for cat, r in finetuned_results.items()}, }, f, indent=2, default=lambda x: None if (isinstance(x, float) and math.isnan(x)) else x) print(f"\nResults saved to: {results_path}") print(f"Details saved to: {details_path}") print(f"\n{'=' * 70}") print("Validation Complete!") print(f"{'=' * 70}") if __name__ == "__main__": main()