hallucination / training /test_nullu.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
2.29 kB
"""Example usage for `load_nullu_model`.
Run Nullu's `experiments/llava_run.pbs` then `experiments/llava_edit.pbs` first to produce
an edited checkpoint (every layer in [0, 32) has its `mlp.down_proj.weight` null-projected).
This script then loads that checkpoint and applies its edits to only the chosen layer slice
on a fresh `HookedSAELlavaConditionalGeneration`, without re-running the edit pipeline.
Example:
python training/test_nullu.py \
--edited-model /path/to/Nullu/output/edited_model/LLaVA-7B-top4-0-32-test \
--lowest-layer 16 --highest-layer 32
"""
import argparse
import torch
from model.llava.hooked_llava import (
HookedSAELlavaConditionalGeneration,
load_nullu_model,
)
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--edited-model",
required=True,
help="Path to Nullu's edited checkpoint directory "
"(e.g. Nullu/output/edited_model/LLaVA-7B-top4-0-32-test).",
)
parser.add_argument("--lowest-layer", type=int, default=16)
parser.add_argument("--highest-layer", type=int, default=32)
parser.add_argument("--base-model", default="llava-hf/llava-1.5-7b-hf")
args = parser.parse_args()
device = "cuda:0" if torch.cuda.is_available() else "cpu"
dtype = torch.float16
model = load_nullu_model(
lowest_layer=args.lowest_layer,
highest_layer=args.highest_layer,
edited_model_path=args.edited_model,
base_model_name=args.base_model,
torch_dtype=dtype,
device=device,
)
base = HookedSAELlavaConditionalGeneration.from_pretrained(
args.base_model, torch_dtype=dtype
).to(device)
n_layers = model.config.text_config.num_hidden_layers
print(f"layer | in-range | status")
for i in range(n_layers):
merged_w = model.model.language_model.layers[i].mlp.down_proj.weight
base_w = base.model.language_model.layers[i].mlp.down_proj.weight
edited = not torch.allclose(merged_w, base_w)
in_range = args.lowest_layer <= i < args.highest_layer
status = "EDITED" if edited else "base"
marker = "*" if in_range else " "
print(f" {i:02d} | {marker} | {status}")
if __name__ == "__main__":
main()