File size: 15,794 Bytes
a2ffd07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
"""
Run knowledge-editing baselines on the hallucination suppression task.

Methods supported:
  - lora:     LoRA fine-tuning via EasyEdit
  - dualedit: DualEdit (vision + text adapters, custom implementation)

Usage:
    python -m experiment.knowledge_editing.run_baselines \
        --edit_set experiment/knowledge_editing/edit_set.json \
        --methods lora dualedit \
        --output_dir step4_ke_outputs
"""

import argparse
import json
import os
import sys
import copy
from datetime import datetime
from pathlib import Path

import torch
from PIL import Image
from tqdm import tqdm

# Ensure project root is on path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))

from experiment.knowledge_editing.llava15_compat import (
    LLaVA15ProcessorWrapper,
    LLaVA15ImageProcessor,
)


def load_edit_set(path: str) -> dict:
    with open(path) as f:
        return json.load(f)


def _load_image(inst: dict, hf_images: dict | None) -> Image.Image | None:
    """Load an image from file path or HF dataset cache."""
    image_path = inst.get("image_path")
    image_id = inst.get("image_id")

    if image_path and os.path.isfile(image_path):
        return Image.open(image_path).convert("RGB")

    if hf_images and image_id in hf_images:
        return hf_images[image_id].convert("RGB")

    return None


def _build_hf_image_cache(dataset_id: str) -> dict:
    """Build an {image_id: PIL.Image} lookup from the HF dataset."""
    from experiment.data.hf_loader import load_hf_dataset

    print(f"  Loading images from HuggingFace dataset ({dataset_id})...")
    ds = load_hf_dataset(dataset_id)
    if hasattr(ds, "keys"):
        from datasets import concatenate_datasets
        ds = concatenate_datasets([ds[s] for s in ds])

    return {item["image_id"]: item["image"] for item in ds}


def build_requests(edit_set: dict, dataset_id: str = "pbcong/bathroom-toilet",
                   use_eval_instances: bool = False):
    """Convert edit_set instances into request format.

    Args:
        use_eval_instances: If True, use eval_instances for the efficacy category
            (HF val split, the fixed 50-image eval set). Used by DualEdit so
            editing and evaluation are on the same images. If False, use
            edit_instances.train (HF train split) for methods like LoRA.
    """
    # Resolve efficacy category name from edit_descriptor
    relation_key = edit_set.get("edit_descriptor", {}).get("relation", "bathroom_toilet")
    efficacy_cat = edit_set.get("edit_descriptor", {}).get("concept", "bathroom_no_toilet")

    if use_eval_instances and "eval_instances" in edit_set:
        instances = edit_set["eval_instances"].get(efficacy_cat, [])
        print(f"  Using eval_instances.{efficacy_cat} ({len(instances)} images, HF val split)")
    else:
        instances = edit_set["edit_instances"]["train"]
        print(f"  Using edit_instances.train ({len(instances)} images, HF train split)")

    locality = edit_set["locality_instances"]

    edit_prompt = edit_set["prompts"]["edit_prompt"]

    generality_prompts = edit_set["prompts"]["generality_prompts"]
    rephrase_prompt = generality_prompts[0] if generality_prompts else edit_prompt

    # Use first available locality category (scene_with_object)
    loc_cat_name = next(iter(locality), None)
    loc_bwt = locality[loc_cat_name] if loc_cat_name else []

    # Filter for usable instances first, then apply n_edits cap so we don't
    # waste the budget on instances that have no target (non-hallucinating images).
    usable_all = [
        inst for inst in instances
        if inst.get("target") is not None and inst.get("is_usable", True)
    ]
    skipped = len(instances) - len(usable_all)
    if skipped:
        print(f"  {skipped} instances skipped (no target or degenerate after cleaning)")
    usable = usable_all

    hf_images = None
    all_instances = list(usable) + list(loc_bwt)
    needs_hf = any(
        not inst.get("image_path") or not os.path.isfile(inst.get("image_path", ""))
        for inst in all_instances
    )
    if needs_hf:
        hf_images = _build_hf_image_cache(dataset_id)

    requests = []
    for i, inst in enumerate(usable):
        edit_image = _load_image(inst, hf_images)
        if edit_image is None:
            print(f"  Skipping {inst['image_id']}: image not found")
            continue

        rephrase_idx = (i + 1) % len(usable)
        rephrase_inst = usable[rephrase_idx]
        rephrase_image = _load_image(rephrase_inst, hf_images) or edit_image

        text_loc_prompt = "What is the capital of France?"
        text_loc_answer = "Paris"

        loc_inst = loc_bwt[i % len(loc_bwt)]
        loc_image = _load_image(loc_inst, hf_images) or edit_image
        loc_vis_prompt = edit_prompt
        loc_vis_answer = loc_inst.get("original_caption") or "A room with various objects."

        request = {
            "prompt": edit_prompt,
            "target": inst["target"],
            "image": edit_image,
            "file_type": "image",
            "rephrase_prompt": rephrase_prompt,
            "image_rephrase": rephrase_image,
            "locality_prompt": text_loc_prompt,
            "locality_ground_truth": text_loc_answer,
            "multimodal_locality_image": loc_image,
            "multimodal_locality_prompt": loc_vis_prompt,
            "multimodal_locality_ground_truth": loc_vis_answer,
            "_image_id": inst["image_id"],
        }
        requests.append(request)

    print(f"Built {len(requests)} edit requests (completion formulation)")
    return requests


def load_model_and_processor(model_name: str, device: str, dtype: torch.dtype):
    """Load LLaVA-1.5 model and wrap processor."""
    from transformers import AutoTokenizer, LlavaForConditionalGeneration, LlavaProcessor
    from transformers import CLIPImageProcessor

    print(f"Loading {model_name}...")
    model = LlavaForConditionalGeneration.from_pretrained(
        model_name, torch_dtype=dtype, device_map={"": device},
    )
    tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=False)
    image_processor = CLIPImageProcessor.from_pretrained(model_name)
    raw_processor = LlavaProcessor(tokenizer=tokenizer, image_processor=image_processor)
    processor = LLaVA15ProcessorWrapper(raw_processor)

    return model, processor


def load_hparams(method: str, hparams_dir: str):
    """Load hparams for a given method."""
    yaml_path = os.path.join(hparams_dir, f"{method}.yaml")
    if not os.path.exists(yaml_path):
        raise FileNotFoundError(f"Hparams not found: {yaml_path}")

    if method == "lora":
        from easyeditor.models.lora import LoRAMultimodalHyperParams
        return LoRAMultimodalHyperParams.from_hparams(yaml_path)
    elif method == "dualedit":
        from experiment.knowledge_editing.dualedit.dualedit_hparams import DualEditHyperParams
        return DualEditHyperParams.from_hparams(yaml_path)
    else:
        raise ValueError(f"Unknown method: {method}")


def get_apply_algo(method: str):
    """Get the algorithm function for a method."""
    if method == "lora":
        from easyeditor.models.lora.lora_main import apply_lora_to_multimodal_model
        return apply_lora_to_multimodal_model
    elif method == "dualedit":
        from experiment.knowledge_editing.dualedit import apply_dualedit_to_multimodal_model
        return apply_dualedit_to_multimodal_model
    else:
        raise ValueError(f"Unknown method: {method}")


def save_edited_model(model, processor, output_dir: str, method: str):
    """Save the edited model for later evaluation."""
    save_dir = os.path.join(output_dir, f"{method}_edited")
    merged_dir = os.path.join(save_dir, "merged_for_eval")
    os.makedirs(merged_dir, exist_ok=True)

    from transformers import PreTrainedModel
    base_model = model
    while not isinstance(base_model, PreTrainedModel) and hasattr(base_model, "model"):
        base_model = base_model.model

    if method == "dualedit":
        if hasattr(base_model, '_dualedit_state'):
            dualedit_state = base_model._dualedit_state
            save_state = {k: v for k, v in dualedit_state.items()
                          if k not in ("vision_hook_handle", "text_hook_handle")}
            adapter_path = os.path.join(merged_dir, "dualedit_state.pt")
            torch.save(save_state, adapter_path)
            print(f"  Saved DualEdit adapter state to {adapter_path}")
        else:
            print("  WARNING: DualEdit state not found on model — nothing saved")
    else:
        # LoRA: edits are merged into the weights
        base_model.save_pretrained(merged_dir)
        processor._processor.save_pretrained(merged_dir)

    print(f"  Saved edited model to {merged_dir}")
    return save_dir


def run_single_method(
    method: str,
    model,
    processor,
    requests: list[dict],
    hparams_dir: str,
    output_dir: str,
):
    """Run one method on the edit requests."""
    print(f"\n{'='*60}")
    print(f"Running: {method.upper()}")
    print(f"  {len(requests)} edit instances")
    print(f"{'='*60}")

    hparams = load_hparams(method, hparams_dir)
    apply_algo = get_apply_algo(method)

    edited_model = model

    try:
        checkpoint_dir = os.path.join(output_dir, "checkpoints")
        edited_model, weights_copy = apply_algo(
            model,
            processor,
            requests,
            hparams,
            copy=False,
            return_orig_weights=True,
            keep_original_weight=False,
            checkpoint_dir=checkpoint_dir,
        )
    except Exception as e:
        print(f"  {method} apply failed: {e}")
        import traceback
        traceback.print_exc()

    save_dir = save_edited_model(edited_model, processor, output_dir, method)

    return edited_model, save_dir


def run_evaluation(
    model_type: str,
    model_dir: str,
    base_model_name: str,
    output_dir: str,
    run_name: str,
    edit_set_path: str,
):
    """Run evaluation pipeline on the edited model."""
    import subprocess

    cmd = [
        sys.executable, "-m", "experiment.evaluation.validate",
        "--model_type", model_type,
        "--model_dir", model_dir,
        "--base_model_name", base_model_name,
        "--inference_backend", "transformers",
        "--mention_method", "keyword",
        "--output_dir", output_dir,
        "--num_per_category", "50",
    ]

    print(f"\n  Running evaluation: {run_name}")
    print(f"  Command: {' '.join(cmd)}")

    result = subprocess.run(cmd, capture_output=True, text=True)
    if result.returncode != 0:
        print(f"  Evaluation failed:\n{result.stderr}")
    else:
        print(f"  Evaluation complete")
        for line in result.stdout.split("\n"):
            if any(k in line for k in ["efficacy", "generality", "locality",
                                        "consistency", "Efficacy", "Generality",
                                        "Locality", "Consistency"]):
                print(f"    {line}")

    return result.returncode == 0


def main():
    parser = argparse.ArgumentParser(
        description="Run LoRA and DualEdit baselines for hallucination suppression"
    )
    parser.add_argument("--edit_set", type=str,
                        default="experiment/knowledge_editing/edit_set.json")
    parser.add_argument("--methods", nargs="+",
                        default=["lora", "dualedit"],
                        choices=["lora", "dualedit"])
    parser.add_argument("--model_name", type=str,
                        default="llava-hf/llava-1.5-7b-hf")
    parser.add_argument("--hparams_dir", type=str,
                        default=os.path.join(os.path.dirname(os.path.abspath(__file__)), "hparams"))
    parser.add_argument("--output_dir", type=str,
                        default="step4_ke_outputs")
    parser.add_argument("--dataset_id", type=str, default=None,
                        help="HuggingFace dataset ID for loading images")
    parser.add_argument("--device", type=str, default="cuda")
    parser.add_argument("--batch", action="store_true",
                        help="Ignored (kept for backwards compatibility)")
    parser.add_argument("--skip_eval", action="store_true",
                        help="Skip evaluation (just run edits and save)")
    args = parser.parse_args()

    run_id = datetime.now().strftime("%Y%m%d_%H%M%S")
    run_dir = os.path.join(args.output_dir, f"ke_run_{run_id}")
    os.makedirs(run_dir, exist_ok=True)

    with open(os.path.join(run_dir, "run_config.json"), "w") as f:
        json.dump(vars(args), f, indent=2)

    print("Loading edit set...")
    edit_set = load_edit_set(args.edit_set)
    print(f"  Stats: {edit_set['stats']}")

    # Resolve dataset_id: CLI > edit_set > default
    dataset_id = args.dataset_id
    if not dataset_id:
        dataset_id = edit_set.get("data_config", {}).get("dataset_id", "pbcong/bathroom-toilet")

    # DualEdit edits on the eval set (HF val, 50 images) so that editing and
    # evaluation use the exact same images. Other methods (LoRA) train on the
    # full HF train split and are evaluated on the separate eval set.
    use_eval = set(args.methods) == {"dualedit"}
    requests = build_requests(edit_set, dataset_id=dataset_id,
                              use_eval_instances=use_eval)
    if not requests:
        print("ERROR: No valid requests built. Check edit_set.json.")
        return

    # Save the exact image IDs used for editing so eval can pin to the same images.
    edit_image_ids = [r["_image_id"] for r in requests]
    edit_image_ids_path = os.path.join(run_dir, "edit_image_ids.json")
    with open(edit_image_ids_path, "w") as f:
        json.dump(edit_image_ids, f)
    print(f"  Saved {len(edit_image_ids)} edit image IDs to {edit_image_ids_path}")

    results_summary = {}

    for method in args.methods:
        print(f"\nLoading fresh model for {method}...")
        model, processor = load_model_and_processor(
            args.model_name, args.device, torch.float16,
        )

        try:
            edited_model, save_dir = run_single_method(
                method=method,
                model=model,
                processor=processor,
                requests=requests,
                hparams_dir=args.hparams_dir,
                output_dir=run_dir,
            )

            results_summary[method] = {
                "status": "edited",
                "save_dir": save_dir,
                "n_edits": len(requests),
            }

            if not args.skip_eval:
                merged_dir = os.path.join(save_dir, "merged_for_eval")
                eval_model_type = "dualedit" if method == "dualedit" else "merged"
                success = run_evaluation(
                    model_type=eval_model_type,
                    model_dir=merged_dir,
                    base_model_name=args.model_name,
                    output_dir=run_dir,
                    run_name=f"{method}_n{len(requests)}",
                    edit_set_path=args.edit_set,
                )
                results_summary[method]["eval_success"] = success

        except Exception as e:
            print(f"  {method} FAILED: {e}")
            import traceback
            traceback.print_exc()
            results_summary[method] = {"status": "failed", "error": str(e)}

    summary_path = os.path.join(run_dir, "results_summary.json")
    with open(summary_path, "w") as f:
        json.dump(results_summary, f, indent=2)

    print(f"\n{'='*60}")
    print("All methods complete.")
    print(f"Results saved to: {run_dir}")
    print(f"Summary: {summary_path}")
    for method, result in results_summary.items():
        print(f"  {method}: {result['status']}")
    print(f"{'='*60}")


if __name__ == "__main__":
    main()