Flatten 1_raw_dataset submodules into plain tracked files

FocalLoRA, Should-It-Be-Executed-Or-Processed, and topicattack were
nested git repos (with an inner FocalLoRA/data/FocalLoRA/.git as well).
Drop their .git history and track the contents directly in this repo
instead of as submodules/gitlinks.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
HenryChou020514
2026-07-07 19:06:09 +08:00
parent 6edf7da2b7
commit 01bb07dba8
167 changed files with 93492 additions and 3 deletions

View File

@ -0,0 +1,691 @@
"""
Focal-Head LoRA Finetune
==========================================
• Selectively fine-tunes "important attention heads" (via LoRA) to enhance LLM alignment with system instructions.
• Key components:
1) detect_heads : compares normal vs. conflict attention → selects top-k heads
2) Q-LoRA (4-bit): injects LoRA only into q/k projection layers with 4-bit quantization
3) make_sys_mask : builds token-level masks for system segments across chat templates
4) focus_loss : encourages final-token attention to return to system region (FP32 for numerical stability)
"""
import os, json, argparse, math, glob, random, re, pickle
from collections import defaultdict
from typing import List, Tuple, Dict
import random
import torch, numpy as np
from torch.utils.data import Dataset, DataLoader
from tqdm import tqdm
from transformers import (
AutoConfig, AutoTokenizer, AutoModelForCausalLM,
BitsAndBytesConfig, get_linear_schedule_with_warmup,
)
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training, PeftModel
import evallib as evallib
# ---------------- Set random seed ----------------
SEED = 42
random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)
# Default evaluation dataset (fixed 8-task dev set)
EVAL_DATA_PATH = os.path.join("../data/focal_lora_dataset_dev/dev_eval.json")
# ======================================================
# 1⃣ Locate LoRA target layers (q_proj/k_proj)
# ======================================================
def get_lora_targets(model, layers: List[int]) -> List[str]:
mtype = (getattr(model.config, "model_type", "") or "").lower()
archs = [x.lower() for x in getattr(model.config, "architectures", [])]
if mtype.startswith("qwen2") or any("qwen2" in a for a in archs):
return [f"model.layers.{i}.self_attn.{p}" for i in layers for p in ("q_proj", "k_proj")]
if "phi" in mtype or any("phi" in a for a in archs):
return [f"model.layers.{i}.self_attn.{p}" for i in layers for p in ("q_proj", "k_proj", "qkv_proj")]
if mtype in {"llama", "mistral"} or "llama" in mtype:
return [f"model.layers.{i}.self_attn.{p}" for i in layers for p in ("q_proj", "k_proj")]
# fallback for unknown models
cand = []
for name, _ in model.named_modules():
if any(f".{i}." in name for i in layers) and name.split(".")[-1] in {
"q_proj", "k_proj", "qkv_proj", "c_attn", "query_key_value"}:
cand.append(name)
return cand
# ======================================================
# 2⃣ Load model with 4-bit quantization
# ======================================================
def load_model(model_path: str):
cfg = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
try:
tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_fast=True)
tok_inf = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_fast=True)
except Exception:
tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_fast=False)
tok_inf = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_fast=False)
if tok.pad_token_id is None:
tok.pad_token = tok.eos_token
tok.pad_token_id = tok.eos_token_id
tok.padding_side = "right"
if tok_inf.pad_token_id is None:
tok_inf.pad_token = tok.eos_token
tok_inf.pad_token_id = tok.eos_token_id
tok_inf.padding_side = "left"
bnb_cfg = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
)
model = AutoModelForCausalLM.from_pretrained(
model_path,
config=cfg,
quantization_config=bnb_cfg,
device_map="auto",
trust_remote_code=True,
attn_implementation="eager",
)
return model, tok, tok_inf
# ======================================================
# 3⃣ Construct system token mask
# ======================================================
def make_sys_mask(input_ids: torch.Tensor, sub_ids: torch.Tensor, tokenizer) -> torch.Tensor:
"""
input_ids: Tensor (B, N)
sub_ids: Tensor (B, M) padded with tokenizer.pad_token_id
tokenizer: tokenizer object with pad_token_id and decode()
Returns:
mask: Bool tensor of shape (B, N) with True only for the FIRST match
"""
pad_id = tokenizer.pad_token_id
device = input_ids.device
B, N = input_ids.shape
_, M = sub_ids.shape
# Compute true (unpadded) lengths
sub_lens = (sub_ids != pad_id).sum(dim=1) # (B,)
mask = torch.zeros_like(input_ids, dtype=torch.bool)
for b in range(B):
L = sub_lens[b].item()
if L == 0 or L > N:
print(f"\n⚠️ Invalid sub length at batch {b}")
print("input_ids:", tokenizer.decode(input_ids[b], skip_special_tokens=False))
print("sub_ids: ", tokenizer.decode(sub_ids[b], skip_special_tokens=False))
continue
# Sliding windows
windows = input_ids[b].unfold(dimension=0, size=L, step=1) # (N-L+1, L)
# Target without padding
target = sub_ids[b, :L] # (L,)
full_match = (windows == target).all(dim=1)
idx = torch.where(full_match)[0]
if len(idx) > 0: # ✅ FIRST match only
start = idx[0].item()
mask[b, start:start + L] = True
else:
# ❌ NOT FOUND → DEBUG OUTPUT
print(f"\n❌ Subsequence NOT found at batch index {b}")
print("input_ids:", tokenizer.decode(input_ids[b], skip_special_tokens=False))
print("sub_ids: ", tokenizer.decode(sub_ids[b, :L], skip_special_tokens=False))
# Attention Sink
B, L = input_ids.shape
non_pad = (input_ids != pad_id) # [B, L], bool
first_nonpad = non_pad.int().argmax(dim=1) # [B]
positions = torch.arange(L, device=input_ids.device).unsqueeze(0) # [1, L]
window_mask = (positions >= first_nonpad.unsqueeze(1)) & \
(positions < (first_nonpad + 4).unsqueeze(1)) & \
non_pad
mask |= window_mask # or: mask = window_mask.clone() if you want only this
return mask
def make_orig_sys_mask(input_ids: torch.Tensor, tok) -> torch.Tensor:
pad_id = tok.pad_token_id
B, L = input_ids.shape
mask = torch.zeros_like(input_ids, dtype=torch.bool)
tid = tok.convert_tokens_to_ids
start_header = tid("<|start_header_id|>")
end_header = tid("<|end_header_id|>")
eot = tok.eos_token_id
sys_tok = tid("<|system|>")
end_tok = tid("<|end|>")
im_start = tid("<|im_start|>")
im_end = tid("<|im_end|>")
inst_start = tid("[INST]")
inst_end = tid("[/INST]")
for b in range(B):
row = input_ids[b].tolist()
# Format a: header template
if start_header in row:
try:
s = row.index(end_header) + 1
e = row.index(eot)
mask[b, s:e] = True
continue
except ValueError:
pass
# Format b: ChatML <|system|>
if sys_tok in row:
try:
s = row.index(sys_tok) + 1
e = row.index(end_tok, s)
mask[b, s:e] = True
continue
except ValueError:
pass
# Format c: OpenChat <|im_start|> system <|im_end|>
if im_start in row and im_end in row:
for pos in [i for i, t in enumerate(row) if t == im_start]:
if pos + 1 < L and tok.decode([row[pos + 1]]).strip() == "system":
s = pos + 2
e = row.index(im_end, s)
mask[b, s:e] = True
break
if mask[b].any():
continue
# Format d: [INST]...[/INST]
if inst_start in row and inst_end in row:
ist = row.index(inst_start) + 1
iend = row.index(inst_end)
split = None
for i in range(ist, iend - 1):
if input_ids[b, i].item() == eot and input_ids[b, i + 1].item() == eot:
split = i
break
if split is None:
for i in range(ist, iend):
if tok.decode([row[i]]).isspace():
split = i
break
if split and ist < split:
mask[b, ist:split] = True
else:
mask[b, ist:iend] = True
# Attention Sink
B, L = input_ids.shape
non_pad = (input_ids != pad_id) # [B, L], bool
first_nonpad = non_pad.int().argmax(dim=1) # [B]
positions = torch.arange(L, device=input_ids.device).unsqueeze(0) # [1, L]
window_mask = (positions >= first_nonpad.unsqueeze(1)) & \
(positions < (first_nonpad + 4).unsqueeze(1)) & \
non_pad
mask |= window_mask # or: mask = window_mask.clone() if you want only this
return mask
# ======================================================
# 4⃣ Identify important attention heads
# ======================================================
def trim_and_stack(rows):
m = min(len(r) for r in rows)
return np.stack([r[:m] for r in rows])
def trim_same(a, b):
m = min(a.shape[1], b.shape[1])
return a[:, :m], b[:, :m]
def score_heads(norm, conf):
scores = {}
for k in norm:
if k not in conf:
continue
try:
n = trim_and_stack(norm[k])
c = trim_and_stack(conf[k])
n, c = trim_same(n, c)
except Exception:
continue
p, q = [np.exp(x - np.max(x, -1, keepdims=True)) for x in (n, c)]
p /= p.sum(-1, keepdims=True)
q /= q.sum(-1, keepdims=True)
kl = (p * (np.log(p + 1e-6) - np.log(q + 1e-6))).sum() / p.shape[0]
shift = np.mean(np.abs(n.mean(1) - c.mean(1)))
frob = np.linalg.norm(n - c, ord="fro")
scores[k] = 0.4 * frob + 0.3 * shift + 0.3 * kl
return scores
def extract_attn(model, tok, sys_msg, usr_msg):
text = tok.apply_chat_template(
[{"role": "system", "content": sys_msg},
{"role": "user", "content": usr_msg}],
tokenize=False, add_generation_prompt=True)
inp = tok(text, return_tensors="pt").to(model.device)
with torch.no_grad():
out = model(**inp, output_attentions=True)
return out.attentions
def detect_heads(json_file, model, tok):
data = json.load(open(json_file, encoding="utf-8"))
grp = defaultdict(lambda: {"normal": None, "conflict": None})
for s in data:
bid = s["id"].replace("_normal", "").replace("_conflict", "")
grp[bid][s["label"]] = s
nA, cA = defaultdict(list), defaultdict(list)
for pair in tqdm(grp.values(), desc="Extract"):
for lab in ("normal", "conflict"):
if pair[lab] is None:
continue
s = pair[lab]
usr = f"{s['task']} {s['user_message']}".strip() or s["task"]
attn = extract_attn(model, tok, s["system_message"], usr)
last = attn[0][0].shape[2] - 1
for l in range(len(attn)):
for h in range(attn[l][0].shape[0]):
row = attn[l][0][h, last, :].float().cpu().numpy()
(nA if lab == "normal" else cA)[f"L{l}_H{h}"].append(row)
scored = sorted(score_heads(nA, cA).items(), key=lambda x: x[1], reverse=True)
return [(k, float(v)) for k, v in scored]
def save_heads_config(heads, output_dir, model_path, json_path, topk):
"""Cache detected heads so we can resume training without recomputing."""
os.makedirs(output_dir, exist_ok=True)
heads_path = os.path.join(output_dir, "heads.json")
payload = {
"model_path": model_path,
"json_path": json_path,
"topk": str(topk),
"heads": heads,
}
with open(heads_path, "w", encoding="utf-8") as f:
json.dump(payload, f, indent=2)
print(f"💾 Saved heads cache → {heads_path}")
def select_top_heads(all_heads: List[Tuple[str, float]], topk_spec) -> List[Tuple[str, float]]:
"""Select top heads based on numeric count or percentage (e.g., '10p')."""
if not all_heads:
return []
if topk_spec is None:
return all_heads
if isinstance(topk_spec, str):
spec = topk_spec.strip().lower()
else:
spec = str(topk_spec)
if not spec:
return all_heads
if spec.endswith("p"):
try:
percent = float(spec[:-1])
except ValueError:
raise ValueError(f"Invalid percentage for --topk: {topk_spec}")
count = max(1, math.ceil(percent / 100.0 * len(all_heads)))
else:
try:
count = int(float(spec))
except ValueError:
raise ValueError(f"Invalid numeric value for --topk: {topk_spec}")
count = max(1, count)
return all_heads[:min(count, len(all_heads))]
def load_heads_config(heads_file: str):
if not os.path.exists(heads_file):
raise FileNotFoundError(f"Heads file not found: {heads_file}")
with open(heads_file, "r", encoding="utf-8") as f:
payload = json.load(f)
raw_heads = payload.get("heads")
if raw_heads is None:
raise ValueError(f"'heads' not defined in {heads_file}")
heads = [(str(tag), float(score)) for tag, score in raw_heads]
meta = {
"model_path": payload.get("model_path"),
"json_path": payload.get("json_path"),
"topk": payload.get("topk"),
}
return heads, meta
def _extract_suffix_index(name: str) -> int:
m = re.search(r"(\d+)$", name)
return int(m.group(1)) if m else -1
def discover_existing_adapter(out_dir: str):
if not os.path.isdir(out_dir):
return None, 0
candidates = []
root_config = os.path.join(out_dir, "adapter_config.json")
if os.path.exists(root_config):
candidates.append((0, out_dir))
for entry in os.listdir(out_dir):
path = os.path.join(out_dir, entry)
if not os.path.isdir(path):
continue
if os.path.exists(os.path.join(path, "adapter_config.json")):
candidates.append((_extract_suffix_index(entry), path))
if not candidates:
return None, 0
candidates.sort(key=lambda x: x[0])
resume_path = candidates[-1][1]
next_idx = candidates[-1][0] + 1 if candidates[-1][0] >= 0 else 0
return resume_path, next_idx
# ======================================================
# 5⃣ Focus Loss: encourages attention to system region
# ======================================================
def focus_loss(attns, sys_mask, heads):
B = sys_mask.size(0)
total_loss = torch.zeros([], dtype=torch.float32, device=sys_mask.device)
valid_heads = 0
for tag, _ in heads:
l = int(tag.split("_")[0][1:])
h = int(tag.split("_H")[1])
A = attns[l][:, h].float()
last = A.size(1) - 1
head_loss = torch.zeros([], dtype=torch.float32, device=sys_mask.device)
for b in range(B):
m = sys_mask[b]
if not m.any():
continue
v = A[b, last]
head_loss += v[m].sum() / v.sum().clamp_min(1e-6) / B
total_loss += head_loss
valid_heads += 1
return 1 - total_loss / max(valid_heads, 1)
# ======================================================
# Dataset and Collate Function for Fine-tuning
# ======================================================
class ConflictDS(Dataset):
"""Dataset for loading conflict samples from multiple JSON files."""
def __init__(self, json_files: List[str], tokenizer):
self.samples = []
self.tokenizer = tokenizer
for json_file in json_files:
if not os.path.exists(json_file):
continue
with open(json_file, 'r', encoding='utf-8') as f:
data = json.load(f)
# Filter for conflict samples only
conflicts = [s for s in data if s.get('label') == 'conflict']
self.samples.extend(conflicts)
random.shuffle(self.samples)
print(f"📊 Loaded {len(self.samples)} conflict samples from {len(json_files)} files")
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
return self.samples[idx]
def collate(batch: List[Dict], tokenizer):
"""
Collate function to batch samples and tokenize them.
Combines task + user_message as described in the paper.
"""
conversations = []
texts_sys = []
for sample in batch:
# Combine task and user_message (if present)
task = sample.get('task', '')
user_msg = sample.get('user_message', '')
# Combine as per line 209 logic: task + user_message
user_content = f"{task} {user_msg}".strip() if user_msg else task
# Build chat format
messages = [
{"role": "system", "content": sample['system_message']},
{"role": "user", "content": user_content}
]
conversations.append(messages)
texts_sys.append(sample['system_message'])
# Apply chat template and tokenize
texts = [
tokenizer.apply_chat_template(conv, tokenize=False, add_generation_prompt=True)
for conv in conversations
]
# Tokenize with padding
encoded = tokenizer(
texts,
padding=True,
truncation=True,
max_length=2048,
return_tensors='pt'
)
encoded_sys = tokenizer(
texts_sys,
padding=True,
truncation=True,
max_length=2048,
add_special_tokens=False,
return_tensors='pt'
)
return {
'input_ids': encoded['input_ids'],
'system_ids': encoded_sys['input_ids'],
'attention_mask': encoded['attention_mask']
}
# ======================================================
# 6⃣ Training with LoRA on selected heads
# ======================================================
def save_model(
model,
tok,
out_dir,
epoch,
batch_idx,
current_ratio,
heads=None,
eval_data_path: str = EVAL_DATA_PATH,
):
"""Save model/tokenizer and run lightweight eval with detailed logging."""
print("Running eval and saving model")
save_dir = os.path.join(out_dir, f"batch_{epoch}_{batch_idx}")
os.makedirs(save_dir, exist_ok=True)
model.save_pretrained(save_dir)
tok.save_pretrained(save_dir)
# Run quick evaluations
eval_asr = evallib.quick_eval_asr(
model,
tokenizer=tok,
data_path=eval_data_path,
heads=heads,
)
eval_mmlu = evallib.quick_eval_mmlu(
model,
tokenizer=tok
)
head_pairs = []
if heads:
for tag, _score in heads:
try:
l = int(tag.split("_")[0][1:])
h = int(tag.split("_")[1][1:])
head_pairs.append((l, h))
except Exception:
continue
# quick_eval_asr handles attention capture internally now; just forward the payload
detail_payload = {"eval_asr": eval_asr, "eval_mmlu": eval_mmlu}
with open(os.path.join(save_dir, "detail_log.pkl"), "wb") as f:
pickle.dump(detail_payload, f)
# Append training log
info_file = os.path.join(out_dir, "training_log.csv")
if not os.path.exists(info_file):
with open(info_file, "w") as info:
info.write("epoch,batch_idx,current_ratio,normal_success,conflict_success,both_success,mmlu_acc\n")
normal_success = eval_asr.get("normal_success") if isinstance(eval_asr, dict) else None
conflict_success = eval_asr.get("conflict_success") if isinstance(eval_asr, dict) else None
both_success = eval_asr.get("both_success") if isinstance(eval_asr, dict) else None
mmlu_acc = eval_mmlu.get("accuracy") if isinstance(eval_mmlu, dict) else None
with open(info_file, "a") as info:
info.write(
f"{epoch},{batch_idx},{current_ratio:.4f},"
f"{normal_success if normal_success is not None else ''},"
f"{conflict_success if conflict_success is not None else ''},"
f"{both_success if both_success is not None else ''},"
f"{mmlu_acc if mmlu_acc is not None else ''}\n"
)
return save_dir, {"asr": eval_asr, "mmlu": eval_mmlu}
def tune(model, tok,tok_inf, heads, data_dir, out_dir, epochs, bs, lr, lam_foc,
resume_adapter=None, start_batch_idx=0):
layers = sorted({int(t.split("_")[0][1:]) for t, _ in heads})
targets = get_lora_targets(model, layers)
if not targets:
raise ValueError("No q/k projection layers found!")
lora_cfg = LoraConfig(r=8, lora_alpha=16, bias="none",
target_modules=targets, task_type="CAUSAL_LM")
model = prepare_model_for_kbit_training(model, use_gradient_checkpointing=False)
if resume_adapter:
if not os.path.exists(resume_adapter):
raise FileNotFoundError(f"LoRA adapter not found: {resume_adapter}")
model = PeftModel.from_pretrained(model, resume_adapter, is_trainable=True)
print(f"♻️ Loaded existing LoRA adapter → {resume_adapter}")
else:
model = get_peft_model(model, lora_cfg)
os.makedirs(out_dir, exist_ok=True)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
ratio = trainable / total * 100
print(f"🔧 Trainable parameters: {trainable:,} / {total:,} ({ratio:.4f}% of total)")
files = glob.glob(os.path.join(data_dir, "*.json"))
dl = DataLoader(ConflictDS(files, tok), batch_size=bs, shuffle=True,
collate_fn=lambda b: collate(b, tok))
opt = torch.optim.AdamW(model.parameters(), lr=lr,)
total = epochs * math.ceil(len(dl))
sch = get_linear_schedule_with_warmup(opt, int(0.05 * total), total)
model.train()
current_idx = start_batch_idx
current_ratio_reached = False
for ep in range(epochs):
if current_ratio_reached:
break
pbar = tqdm(enumerate(dl), desc=f"Epoch {ep+1}/{epochs}")
for idx,batch in pbar:
batch = {k: v.to(model.device) for k, v in batch.items()}
#breakpoint()
out = model(**batch, output_attentions=True)
# sys_mask = make_sys_mask(batch["input_ids"], batch["system_ids"],tok)
sys_mask = make_orig_sys_mask(batch["input_ids"], tok)
loss = lam_foc * focus_loss(out.attentions, sys_mask, heads)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
sch.step()
opt.zero_grad()
pbar.set_postfix(loss=f"{loss.item():.4f}")
current_ratio = float(loss.detach().cpu().item()) / lam_foc
current_ratio = 1 - current_ratio
if idx % 100 == 0:
save_path, eval_summary = save_model( model, tok_inf, out_dir, current_idx, idx, current_ratio, heads=heads )
eval_metrics = eval_summary.get("asr", {}) if isinstance(eval_summary, dict) else {}
conflict_success = eval_metrics.get("conflict_success")
print(f"✅ LoRA adapter checkpoint saved → {save_path} (conflict_success={conflict_success if conflict_success is not None else 'n/a'})")
save_path, eval_summary = save_model(
model, tok_inf, out_dir, current_idx, idx, current_ratio, heads=heads
)
eval_metrics = eval_summary.get("asr", {}) if isinstance(eval_summary, dict) else {}
conflict_success = eval_metrics.get("conflict_success")
print(f"✅ LoRA adapter saved → {save_path} (conflict_success={conflict_success if conflict_success is not None else 'n/a'})")
current_idx += 1
def main():
ap = argparse.ArgumentParser("Important-Head LoRA Finetune")
ap.add_argument("--json_path", required=False, help="Probing file with normal and conflict samples")
ap.add_argument("--model_path", required=False, help="Base model path")
ap.add_argument("--tune_path", required=True, help="Folder with conflict samples for fine-tuning")
ap.add_argument("--output_dir", default="outputs_lora", help="Path to save LoRA adapter")
ap.add_argument("--lora_path", default=None, help="Optional existing LoRA adapter to load before training")
ap.add_argument("--topk", type=str, default="10",
help="Top-K important heads to select (e.g., 10 or 10p for 10%)")
ap.add_argument("--epochs", type=int, default=3)
ap.add_argument("--batch_size", type=int, default=4)
ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--lambda_focus", type=float, default=0.5)
ap.add_argument("--head_path", type=str, default="", help="Optional path to a precomputed heads.json file.")
args = ap.parse_args()
preferred_heads = args.head_path.strip()
heads_file = preferred_heads or os.path.join(args.output_dir, "heads.json")
heads_meta = {}
all_heads = None
if heads_file and os.path.exists(heads_file):
all_heads, heads_meta = load_heads_config(heads_file)
print(f"📂 Loaded cached heads from {heads_file}")
else:
if preferred_heads:
ap.error(f"--head_path specified but not found: {heads_file}")
if not args.json_path:
ap.error("--json_path is required when no cached heads are found.")
if not args.model_path:
ap.error("--model_path is required when computing new heads.")
model_path = args.model_path or heads_meta.get("model_path")
if not model_path:
ap.error("Base model path missing. Provide --model_path or ensure model_path exists in output_dir/heads.json")
if args.lora_path:
if not os.path.exists(args.lora_path):
ap.error(f"--lora_path not found: {args.lora_path}")
resume_adapter, start_idx = args.lora_path, 0
print(f"♻️ Loaded LoRA adapter from --lora_path: {resume_adapter}")
else:
resume_adapter, start_idx = discover_existing_adapter(args.output_dir)
if resume_adapter:
print(f"♻️ Resuming from existing adapter in output_dir: {resume_adapter}")
model, tok ,tok_inf= load_model(model_path)
if all_heads is not None:
print("📌 Important heads:", all_heads)
else:
all_heads = detect_heads(args.json_path, model, tok)
print("📌 Important heads:", all_heads)
save_heads_config(all_heads, args.output_dir, model_path, args.json_path, args.topk)
heads = select_top_heads(all_heads, args.topk)
print(f"🎯 Using {len(heads)} heads based on topk={args.topk}: {heads}")
tune(model, tok,tok_inf, heads,
args.tune_path, args.output_dir,
args.epochs, args.batch_size, args.lr, args.lambda_focus,
resume_adapter=resume_adapter, start_batch_idx=start_idx)
if __name__ == "__main__":
main()