Pipeline (stages 1 -> 4-1) can now be run in order from each stage folder.
Stage scripts:
- 2-1: make SEP/FocalLora prep portable (derive paths from __file__ instead of
hardcoded /home/hujk/...) and add prepare_head_ident_dataset.sh runner.
Verified the SEP converter reproduces the committed jsonl byte-for-byte.
- 2-2: unify the four Ident_IH_ALL_1-4_<model>.sh scripts (modernise llama to
conda hook + $ROOT/models; add the missing FocalLora step to qwen3-4b/8b so
focallora.json gets generated for them too).
- 2-3: default TARGETS now covers the three curves from the README
(all_roc_inst_0.1, user_roc_inst_0.1, focallora).
- 3-2: add combos/ with 24 scripts (4 models x {pbs,nts,nts_wam} x {squad,tri}),
head ranking pinned to all_roc_inst_0.1, TOPK overridable.
- 4-1: add eval_single.sh driver + combos/ with 24 cross-eval wrappers
(squad-trained -> tri-eval and vice versa), reusing the --eval-only path.
Eval semantics:
- Judge ASR before UTIL: a response carrying the injected answer now counts as
attacked even when it also contains the correct answer. This changes the
metric, so old training_log.csv rows are not comparable.
- Add --dev-holdout: reserve the last N source rows as a dev slice; training
drops them and the in-training quick eval uses only them. Previously the
quick eval silently defaulted to the squad evaluation set, which contradicted
the README and self-contaminated squad-trained runs.
- train_attn_kl_clean.sh now passes --eval-data-path/--eval-topicattack-path.
- Add --eval-step0 to log an untuned-baseline row before any weight update.
Housekeeping:
- Quarantine superseded entry points under legacy/ (2-2 single-step wrappers,
3-2 old _tuning.fix.* wrappers, 3-1 auxiliary), each with a README.
- Fix .gitignore: the model_score rule was anchored at the repo root and never
matched Codes/..., so ~26GB of intermediates had been staged. Now excludes
*.pkl (~25GB), heads_sorted_eval/ (~690MB), outputs_lora/ checkpoints
(~3.2GB) and pycache. heads_sorted/ and head_scoring_combined.json are kept
deliberately: they are small and are the HEAD_PATH inputs stage 3-2 needs.
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
833 lines
35 KiB
Python
Executable File
833 lines
35 KiB
Python
Executable File
"""
|
|
Head-LoRA Finetune (Refactored v2)
|
|
==========================================
|
|
Target: Align specific attention heads to ignore 'data' tokens using Flash Attention Interception.
|
|
|
|
Key features:
|
|
- LoRA head masking on Q proj: only selected heads' dimensions are updated.
|
|
- LoRA head masking on K proj: only KV groups corresponding to selected heads
|
|
are updated (GQA-aware).
|
|
- Preservation loss: heads that share a KV group with selected heads but are
|
|
NOT themselves selected get a KL loss to stay close to the original model,
|
|
since their K changed but they weren't intended training targets.
|
|
- Heads whose KV group is completely untouched need no loss at all.
|
|
"""
|
|
|
|
import os
|
|
import json
|
|
import csv
|
|
import argparse
|
|
import math
|
|
import random
|
|
import re
|
|
import sys
|
|
from typing import List, Tuple, Dict, Set
|
|
|
|
import torch
|
|
import numpy as np
|
|
from tqdm import tqdm
|
|
from torch.utils.data import Dataset, DataLoader
|
|
from evallib import quick_eval_mmlu, quick_eval_asr_util
|
|
from transformers import (
|
|
AutoConfig, AutoTokenizer, AutoModelForCausalLM,
|
|
BitsAndBytesConfig, get_linear_schedule_with_warmup,
|
|
)
|
|
from transformers.models.llama.modeling_llama import ALL_ATTENTION_FUNCTIONS
|
|
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training, PeftModel
|
|
from peft.tuners.lora.layer import Linear as LoraLinear
|
|
|
|
proj_path = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
sys.path.insert(0, proj_path)
|
|
|
|
from lib_code.tokenize_data_mask import apply_chat_tokenize_with_strip_and_mark
|
|
from lib_code.printcolor import print_tok_color, print_data_color_in_batch
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 1. Flash Attention Spy Interceptor
|
|
# -----------------------------------------------------------------------------
|
|
print("🥷 Injecting Spy Interceptor into Flash Attention 2...")
|
|
|
|
if "flash_attention_2" not in ALL_ATTENTION_FUNCTIONS:
|
|
raise RuntimeError("Current environment does not support Flash Attention 2!")
|
|
|
|
orig_flash_attn = ALL_ATTENTION_FUNCTIONS["flash_attention_2"]
|
|
|
|
def _kv_indices_for_heads(heads: List[int], num_q_heads: int, num_kv_heads: int) -> List[int]:
|
|
group_size = num_q_heads // num_kv_heads
|
|
return [h // group_size for h in heads]
|
|
|
|
|
|
def _apply_attention_mask_chunk(logits, attention_mask, start: int, end: int):
|
|
if attention_mask is None:
|
|
return logits
|
|
|
|
if attention_mask.dim() == 4:
|
|
if attention_mask.size(2) > 1:
|
|
mask_slice = attention_mask[..., -1:, start:end]
|
|
else:
|
|
mask_slice = attention_mask[..., :, start:end]
|
|
return logits + mask_slice.squeeze(2)
|
|
|
|
if attention_mask.dim() == 2:
|
|
mask_slice = attention_mask[:, start:end]
|
|
if attention_mask.dtype == torch.bool or (attention_mask.max() <= 1.0 and attention_mask.min() >= 0.0):
|
|
min_dtype = torch.finfo(logits.dtype).min
|
|
return logits.masked_fill(mask_slice[:, None, :] == 0, min_dtype)
|
|
return logits + mask_slice[:, None, :]
|
|
|
|
return logits
|
|
|
|
|
|
def _last_token_logits_chunk(query, key, heads: List[int], start: int, end: int, scaling):
|
|
num_q_heads = query.shape[1]
|
|
num_kv_heads = key.shape[1]
|
|
kv_indices = _kv_indices_for_heads(heads, num_q_heads, num_kv_heads)
|
|
q = query[:, heads, -1, :].float()
|
|
k = key[:, kv_indices, start:end, :].float()
|
|
return (q[:, :, None, :] * k).sum(dim=-1) * scaling
|
|
|
|
|
|
def _chunked_softmax_stats(query, key, attention_mask, heads: List[int], scaling, chunk_size: int):
|
|
seq_len = key.shape[2]
|
|
max_scores = None
|
|
|
|
for start in range(0, seq_len, chunk_size):
|
|
end = min(start + chunk_size, seq_len)
|
|
logits = _last_token_logits_chunk(query, key, heads, start, end, scaling)
|
|
logits = _apply_attention_mask_chunk(logits, attention_mask, start, end)
|
|
cur_max = logits.max(dim=-1).values
|
|
max_scores = cur_max if max_scores is None else torch.maximum(max_scores, cur_max)
|
|
|
|
denom = torch.zeros_like(max_scores, dtype=torch.float32)
|
|
for start in range(0, seq_len, chunk_size):
|
|
end = min(start + chunk_size, seq_len)
|
|
logits = _last_token_logits_chunk(query, key, heads, start, end, scaling)
|
|
logits = _apply_attention_mask_chunk(logits, attention_mask, start, end)
|
|
denom = denom + torch.exp(logits - max_scores[:, :, None]).sum(dim=-1)
|
|
|
|
return max_scores, denom.clamp_min(1e-20)
|
|
|
|
|
|
def _store_teacher_attention(module, query, key, attention_mask, scaling, target_heads_dict, chunk_size: int):
|
|
heads = list(target_heads_dict.keys())
|
|
if not heads:
|
|
return
|
|
|
|
seq_len = key.shape[2]
|
|
max_scores, denom = _chunked_softmax_stats(query, key, attention_mask, heads, scaling, chunk_size)
|
|
chunks = []
|
|
for start in range(0, seq_len, chunk_size):
|
|
end = min(start + chunk_size, seq_len)
|
|
logits = _last_token_logits_chunk(query, key, heads, start, end, scaling)
|
|
logits = _apply_attention_mask_chunk(logits, attention_mask, start, end)
|
|
probs = torch.exp(logits - max_scores[:, :, None]) / denom[:, :, None]
|
|
chunks.append(probs.detach().to(dtype=torch.float16, device="cpu"))
|
|
|
|
probs_all = torch.cat(chunks, dim=-1)
|
|
for idx, h_idx in enumerate(heads):
|
|
target_heads_dict[h_idx] = probs_all[:, idx : idx + 1, :]
|
|
|
|
|
|
def _teacher_stack(teacher_map, layer_idx: int, heads: List[int]):
|
|
return torch.cat(
|
|
[teacher_map[layer_idx][h].to(dtype=torch.float32) for h in heads],
|
|
dim=1,
|
|
)
|
|
|
|
|
|
def _teacher_const(teacher, eps: float):
|
|
return (teacher * torch.log(teacher.clamp_min(eps))).sum(dim=-1)
|
|
|
|
|
|
def _append_selected_chunked_loss(module, query, key, attention_mask, scaling, state, heads: List[int], chunk_size: int):
|
|
eps = state.get("eps", 1e-8)
|
|
layer_idx = module.layer_idx
|
|
device = query.device
|
|
seq_len = key.shape[2]
|
|
data_mask = state["data_mask"].to(device=device, dtype=torch.bool)
|
|
valid_mask = ~data_mask
|
|
teacher = _teacher_stack(state["base_map"], layer_idx, heads)
|
|
teacher_sum = teacher.sum(dim=-1).to(device=device)
|
|
const = _teacher_const(teacher, eps).to(device=device)
|
|
|
|
max_scores, denom = _chunked_softmax_stats(query, key, attention_mask, heads, scaling, chunk_size)
|
|
data_mass = torch.zeros_like(max_scores, dtype=torch.float32)
|
|
valid_mass = torch.zeros_like(max_scores, dtype=torch.float32)
|
|
cross_log_prob = torch.zeros_like(max_scores, dtype=torch.float32)
|
|
|
|
for start in range(0, seq_len, chunk_size):
|
|
end = min(start + chunk_size, seq_len)
|
|
logits = _last_token_logits_chunk(query, key, heads, start, end, scaling)
|
|
logits = _apply_attention_mask_chunk(logits, attention_mask, start, end)
|
|
probs = torch.exp(logits - max_scores[:, :, None]) / denom[:, :, None]
|
|
teacher_chunk = teacher[:, :, start:end].to(device=device)
|
|
data_chunk = data_mask[:, None, start:end]
|
|
valid_chunk = valid_mask[:, None, start:end]
|
|
data_mass = data_mass + (probs * data_chunk.float()).sum(dim=-1)
|
|
valid_mass = valid_mass + (probs * valid_chunk.float()).sum(dim=-1)
|
|
cross_log_prob = cross_log_prob + (teacher_chunk * torch.log(probs.clamp_min(eps))).sum(dim=-1)
|
|
|
|
cross_log_tuned_valid = cross_log_prob - torch.log(valid_mass.clamp_min(eps)) * teacher_sum
|
|
kl = const - cross_log_tuned_valid
|
|
state["selected_losses"].append((kl + state.get("lambda_data", 1.0) * data_mass).mean())
|
|
|
|
|
|
def _append_preservation_chunked_loss(module, query, key, attention_mask, scaling, state, heads: List[int], chunk_size: int):
|
|
eps = state.get("eps", 1e-8)
|
|
layer_idx = module.layer_idx
|
|
device = query.device
|
|
seq_len = key.shape[2]
|
|
teacher = _teacher_stack(state["orig_map"], layer_idx, heads)
|
|
const = _teacher_const(teacher, eps).to(device=device)
|
|
|
|
max_scores, denom = _chunked_softmax_stats(query, key, attention_mask, heads, scaling, chunk_size)
|
|
cross_log_prob = torch.zeros_like(max_scores, dtype=torch.float32)
|
|
|
|
for start in range(0, seq_len, chunk_size):
|
|
end = min(start + chunk_size, seq_len)
|
|
logits = _last_token_logits_chunk(query, key, heads, start, end, scaling)
|
|
logits = _apply_attention_mask_chunk(logits, attention_mask, start, end)
|
|
probs = torch.exp(logits - max_scores[:, :, None]) / denom[:, :, None]
|
|
teacher_chunk = teacher[:, :, start:end].to(device=device)
|
|
cross_log_prob = cross_log_prob + (teacher_chunk * torch.log(probs.clamp_min(eps))).sum(dim=-1)
|
|
|
|
state["preserve_losses"].append((const - cross_log_prob).mean())
|
|
|
|
|
|
def wrapped_flash_attn(module, query, key, value, attention_mask, scaling, **kwargs):
|
|
chunk_size = int(getattr(module.config, "retrieve_attn_chunk_size", 4096))
|
|
|
|
retrieve_map = getattr(module.config, "retrieve_attn_map", None)
|
|
if retrieve_map is not None and hasattr(module, "layer_idx") and module.layer_idx in retrieve_map:
|
|
_store_teacher_attention(
|
|
module,
|
|
query,
|
|
key,
|
|
attention_mask,
|
|
scaling,
|
|
retrieve_map[module.layer_idx],
|
|
chunk_size,
|
|
)
|
|
|
|
loss_state = getattr(module.config, "retrieve_attn_loss_state", None)
|
|
if loss_state is not None and hasattr(module, "layer_idx"):
|
|
layer_idx = module.layer_idx
|
|
selected_heads = loss_state.get("selected", {}).get(layer_idx, [])
|
|
if selected_heads:
|
|
_append_selected_chunked_loss(
|
|
module, query, key, attention_mask, scaling, loss_state, selected_heads, chunk_size
|
|
)
|
|
affected_heads = loss_state.get("affected", {}).get(layer_idx, [])
|
|
if affected_heads:
|
|
_append_preservation_chunked_loss(
|
|
module, query, key, attention_mask, scaling, loss_state, affected_heads, chunk_size
|
|
)
|
|
|
|
orig_result = orig_flash_attn(module, query, key, value, attention_mask, scaling=scaling, **kwargs)
|
|
return orig_result
|
|
|
|
ALL_ATTENTION_FUNCTIONS["flash_attention_2"] = wrapped_flash_attn
|
|
print("✅ Spy Interceptor installed successfully!")
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 2. Utils: Model & Heads
|
|
# -----------------------------------------------------------------------------
|
|
SEED = 42
|
|
random.seed(SEED)
|
|
np.random.seed(SEED)
|
|
torch.manual_seed(SEED)
|
|
|
|
def get_lora_targets(model, layers: List[int]) -> List[str]:
|
|
mtype = (getattr(model.config, "model_type", "") or "").lower()
|
|
if "llama" in mtype or "mistral" in mtype:
|
|
return [f"model.layers.{i}.self_attn.{p}" for i in layers for p in ("q_proj", "k_proj")]
|
|
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"}:
|
|
cand.append(name)
|
|
return cand
|
|
|
|
def load_model(model_path: str):
|
|
cfg = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
|
|
tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True, use_fast=True, padding_side="left")
|
|
|
|
if tok.pad_token_id is None:
|
|
tok.pad_token = tok.eos_token
|
|
tok.pad_token_id = tok.eos_token_id
|
|
|
|
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="flash_attention_2",
|
|
)
|
|
return model, tok, tok
|
|
|
|
def get_target_structure(heads_file: str, topk_spec: str) -> Dict[int, List[int]]:
|
|
if not os.path.exists(heads_file):
|
|
raise FileNotFoundError(f"Heads file not found: {heads_file}")
|
|
|
|
with open(heads_file, "r") as f:
|
|
ihead_list = json.load(f)
|
|
if type(ihead_list[0]) == str:
|
|
all_heads = [(str(tag), 0.0) for tag in ihead_list]
|
|
else:
|
|
all_heads = ihead_list
|
|
|
|
count = len(all_heads)
|
|
if topk_spec.endswith("p"):
|
|
count = max(1, math.ceil(float(topk_spec[:-1]) / 100.0 * len(all_heads)))
|
|
else:
|
|
count = int(topk_spec)
|
|
|
|
target_heads = all_heads[:min(count, len(all_heads))]
|
|
|
|
structure = {}
|
|
for tag, _ in target_heads:
|
|
try:
|
|
parts = tag.split('_')
|
|
l = int(parts[0][1:])
|
|
h = int(parts[1][1:])
|
|
structure.setdefault(l, []).append(h)
|
|
except Exception:
|
|
continue
|
|
|
|
print(f"🎯 Selected {len(target_heads)} heads from {topk_spec}.")
|
|
return structure
|
|
|
|
def discover_existing_adapter(out_dir: str):
|
|
if not os.path.isdir(out_dir): return None, 0, 0
|
|
candidates = []
|
|
if os.path.exists(os.path.join(out_dir, "adapter_config.json")):
|
|
candidates.append((0, 0, out_dir))
|
|
|
|
for entry in os.listdir(out_dir):
|
|
path = os.path.join(out_dir, entry)
|
|
if os.path.isdir(path) and os.path.exists(os.path.join(path, "adapter_config.json")):
|
|
m = re.match(r"^batch_(\d+)_(\d+)$", entry)
|
|
if m:
|
|
epoch_idx = int(m.group(1))
|
|
batch_idx = int(m.group(2))
|
|
candidates.append((epoch_idx, batch_idx, path))
|
|
|
|
if not candidates: return None, 0, 0
|
|
candidates.sort(key=lambda x: (x[0], x[1]))
|
|
return candidates[-1][2], candidates[-1][0], candidates[-1][1]
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 3. GQA-Aware Head Analysis
|
|
# -----------------------------------------------------------------------------
|
|
def compute_gqa_affected_heads(target_structure: Dict[int, List[int]],
|
|
num_q_heads: int,
|
|
num_kv_heads: int) -> Dict[int, List[int]]:
|
|
"""
|
|
For each layer, find Q heads that are NOT selected but share a KV group
|
|
with at least one selected head. These heads will have their K changed
|
|
by LoRA but were not intended as training targets.
|
|
|
|
Returns: {layer_idx: [affected_but_unselected_head_indices]}
|
|
"""
|
|
group_size = num_q_heads // num_kv_heads
|
|
|
|
affected_structure = {}
|
|
for l, selected_heads in target_structure.items():
|
|
selected_set = set(selected_heads)
|
|
|
|
# Find which KV groups are touched
|
|
touched_kv_groups = set()
|
|
for h in selected_heads:
|
|
touched_kv_groups.add(h // group_size)
|
|
|
|
# Find unselected heads in those touched groups
|
|
affected = []
|
|
for kv_g in touched_kv_groups:
|
|
for offset in range(group_size):
|
|
q_head = kv_g * group_size + offset
|
|
if q_head not in selected_set:
|
|
affected.append(q_head)
|
|
|
|
if affected:
|
|
affected_structure[l] = affected
|
|
|
|
return affected_structure
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 4. LoRA Head Masking (Monkey-Patch)
|
|
# -----------------------------------------------------------------------------
|
|
def apply_head_mask_to_lora(model, target_structure: Dict[int, List[int]],
|
|
num_q_heads: int, num_kv_heads: int, head_dim: int):
|
|
"""
|
|
Monkey-patch LoRA layers so that:
|
|
- q_proj: LoRA output is zeroed for unselected Q heads
|
|
- k_proj: LoRA output is zeroed for KV groups not associated with any selected head
|
|
"""
|
|
group_size = num_q_heads // num_kv_heads
|
|
|
|
for name, module in model.named_modules():
|
|
if not isinstance(module, LoraLinear):
|
|
continue
|
|
|
|
# Parse layer index from name
|
|
layer_idx = None
|
|
parts = name.split(".")
|
|
for i, part in enumerate(parts):
|
|
if part == "layers" and i + 1 < len(parts) and parts[i + 1].isdigit():
|
|
layer_idx = int(parts[i + 1])
|
|
break
|
|
|
|
if layer_idx is None or layer_idx not in target_structure:
|
|
continue
|
|
|
|
selected_heads = set(target_structure[layer_idx])
|
|
is_q = name.endswith("q_proj")
|
|
is_k = name.endswith("k_proj")
|
|
|
|
if not (is_q or is_k):
|
|
continue
|
|
|
|
if is_q:
|
|
mask = torch.zeros(num_q_heads * head_dim)
|
|
for h in selected_heads:
|
|
mask[h * head_dim : (h + 1) * head_dim] = 1.0
|
|
desc = f"Q mask: {len(selected_heads)}/{num_q_heads} heads"
|
|
else:
|
|
touched_kv_groups = set()
|
|
for h in selected_heads:
|
|
touched_kv_groups.add(h // group_size)
|
|
|
|
mask = torch.zeros(num_kv_heads * head_dim)
|
|
for g in touched_kv_groups:
|
|
mask[g * head_dim : (g + 1) * head_dim] = 1.0
|
|
desc = f"K mask: {len(touched_kv_groups)}/{num_kv_heads} KV groups"
|
|
|
|
module.register_buffer("head_mask", mask)
|
|
|
|
orig_forward = module.forward
|
|
|
|
def make_masked_forward(orig_fn, mod):
|
|
def masked_forward(x, *args, **kwargs):
|
|
result = orig_fn(x, *args, **kwargs)
|
|
base_out = torch.nn.functional.linear(x, mod.base_layer.weight, mod.base_layer.bias)
|
|
lora_delta = result - base_out
|
|
masked_delta = lora_delta * mod.head_mask.to(lora_delta.device)
|
|
return base_out + masked_delta
|
|
return masked_forward
|
|
|
|
module.forward = make_masked_forward(orig_forward, module)
|
|
print(f" 🎭 Patched {name}: {desc}")
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 5. Loss Functions
|
|
# -----------------------------------------------------------------------------
|
|
def head_attention_loss(tuned_map, base_map, data_mask, lambda_data=1.0, eps=1e-8):
|
|
"""Loss for SELECTED heads: KL toward blind teacher + data mass penalty."""
|
|
mask_data = data_mask.bool()
|
|
mask_valid = ~mask_data
|
|
|
|
if mask_valid.dim() == 2:
|
|
mask_valid = mask_valid.unsqueeze(1)
|
|
mask_data = mask_data.unsqueeze(1)
|
|
|
|
total = 0.0
|
|
n = 0
|
|
|
|
for l, heads in tuned_map.items():
|
|
for h, tuned in heads.items():
|
|
base = base_map.get(l, {}).get(h, None)
|
|
if tuned is None or base is None:
|
|
continue
|
|
|
|
tuned = tuned.float()
|
|
base = base.float().to(tuned.device)
|
|
|
|
base_v = base * mask_valid.float()
|
|
base_v = base_v / base_v.sum(dim=-1, keepdim=True).clamp_min(eps)
|
|
|
|
tuned_v = tuned * mask_valid.float()
|
|
tuned_v_sum = tuned_v.sum(dim=-1, keepdim=True).clamp_min(eps)
|
|
tuned_v = tuned_v / tuned_v_sum
|
|
|
|
kl = (base_v * (torch.log(base_v + eps) - torch.log(tuned_v + eps))).sum(dim=-1).mean()
|
|
data_mass = (tuned * mask_data.float()).sum(dim=-1).mean()
|
|
|
|
total = total + kl + lambda_data * data_mass
|
|
n += 1
|
|
|
|
return total / max(n, 1)
|
|
|
|
|
|
def head_preservation_loss(tuned_map, orig_map, eps=1e-8):
|
|
"""Loss for KV-group-affected but unselected heads: KL toward original model."""
|
|
total = 0.0
|
|
n = 0
|
|
for l, heads in orig_map.items():
|
|
for h, orig in heads.items():
|
|
tuned = tuned_map.get(l, {}).get(h, None)
|
|
if tuned is None or orig is None:
|
|
continue
|
|
tuned = tuned.float()
|
|
orig = orig.float().to(tuned.device)
|
|
kl = (orig * (torch.log(orig + eps) - torch.log(tuned + eps))).sum(dim=-1).mean()
|
|
total += kl
|
|
n += 1
|
|
return total / max(n, 1)
|
|
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 6. Data Loading
|
|
# -----------------------------------------------------------------------------
|
|
class JsonlMessagesDS(Dataset):
|
|
def __init__(self, data_path: str, repeat_single_sample: int = 1):
|
|
self.samples = []
|
|
with open(data_path, "r", encoding="utf-8") as f:
|
|
if data_path.endswith(".json"):
|
|
payload = json.load(f)
|
|
if not isinstance(payload, list):
|
|
raise ValueError(f"Training JSON must contain a list: {data_path}")
|
|
for obj in payload:
|
|
self.samples.append(obj["messages"] if isinstance(obj, dict) and "messages" in obj else obj)
|
|
else:
|
|
for line in f:
|
|
if line.strip():
|
|
obj = json.loads(line)
|
|
self.samples.append(obj["messages"] if isinstance(obj, dict) and "messages" in obj else obj)
|
|
if repeat_single_sample > 1:
|
|
if len(self.samples) != 1:
|
|
raise ValueError("--repeat-single-sample requires a dataset containing exactly one sample.")
|
|
self.samples = self.samples * repeat_single_sample
|
|
random.shuffle(self.samples)
|
|
print(f"📊 Loaded {len(self.samples)} samples.")
|
|
|
|
def __len__(self): return len(self.samples)
|
|
def __getitem__(self, idx): return self.samples[idx]
|
|
|
|
def collate(batch, tokenizer):
|
|
input_ids, attention_mask, data_mask = apply_chat_tokenize_with_strip_and_mark(
|
|
batch, tokenizer, device="cpu", add_generation_prompt=True,
|
|
encode_kwargs={"padding_side": "left"},
|
|
mode="custom_mask == 'inst'",
|
|
custom_mask_identifier={"data": ["<data>", "</data>"], "inst": ["<inst>", "</inst>"]},
|
|
return_tensors="pt"
|
|
)
|
|
return {"input_ids": input_ids, "attention_mask": attention_mask, "data_mask": data_mask}
|
|
|
|
|
|
def forward_for_attention_loss(model, input_ids, attention_mask, logits_to_keep=1):
|
|
kwargs = {
|
|
"input_ids": input_ids,
|
|
"attention_mask": attention_mask,
|
|
"output_attentions": False,
|
|
"use_cache": False,
|
|
}
|
|
if logits_to_keep is not None and logits_to_keep >= 0:
|
|
kwargs["logits_to_keep"] = logits_to_keep
|
|
try:
|
|
return model(**kwargs)
|
|
except TypeError as exc:
|
|
if "logits_to_keep" not in str(exc):
|
|
raise
|
|
kwargs.pop("logits_to_keep", None)
|
|
return model(**kwargs)
|
|
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 7. Training Loop
|
|
# -----------------------------------------------------------------------------
|
|
def save_model(model, tok, out_dir, epoch, batch_idx, data_path, skip_eval=False):
|
|
print("Saving checkpoint...")
|
|
|
|
temp_map = getattr(model.config, "retrieve_attn_map", None)
|
|
if temp_map is not None:
|
|
del model.config.retrieve_attn_map
|
|
loss_state = getattr(model.config, "retrieve_attn_loss_state", None)
|
|
if loss_state is not None:
|
|
del model.config.retrieve_attn_loss_state
|
|
|
|
save_dir = os.path.join(out_dir, f"batch_{epoch}_{batch_idx}")
|
|
model.save_pretrained(save_dir)
|
|
tok.save_pretrained(save_dir)
|
|
|
|
print(f"✅ Saved to {save_dir}")
|
|
|
|
if skip_eval:
|
|
print("⏭️ Skipped checkpoint evaluation.")
|
|
return save_dir
|
|
|
|
mllu_result = quick_eval_mmlu(model, tok)
|
|
asr_result = quick_eval_asr_util(
|
|
model, tok, training_data_path=data_path,
|
|
)
|
|
log_path = os.path.join(out_dir, "training_log.csv")
|
|
write_header = not os.path.exists(log_path)
|
|
with open(log_path, "a", newline="") as f:
|
|
writer = csv.writer(f)
|
|
attack_columns = [
|
|
("naive", "a_naive", "v_naive"),
|
|
("ignore", "a_ignore", "v_ignore"),
|
|
("escape_separation", "a_escape", "v_escape"),
|
|
("completion_realcmb", "a_cmb", "v_cmb"),
|
|
("conv_attack", "a_conv", "v_conv"),
|
|
("none", "a_none", "v_none"),
|
|
]
|
|
if write_header:
|
|
header = ["epoch", "batch", "mmlu score", "asr_mode", "asr_status"]
|
|
for _, a_col, v_col in attack_columns:
|
|
header.append(a_col)
|
|
header.append(v_col)
|
|
writer.writerow(header)
|
|
row = [epoch, batch_idx, mllu_result.get("accuracy"), asr_result.get("eval_mode"), asr_result.get("status")]
|
|
metrics = asr_result.get("metrics", {}) if asr_result.get("status") == "ok" else {}
|
|
for attack, _, _ in attack_columns:
|
|
values = metrics.get(attack, {})
|
|
row.append(values.get("asr"))
|
|
row.append(values.get("valid_rate"))
|
|
writer.writerow(row)
|
|
return save_dir
|
|
|
|
|
|
def tune_new(model, tok, tok_inf, target_structure, data_path, out_dir, epochs, bs, lr,
|
|
resume_adapter=None, start_epoch=0, start_batch_idx=0, batch_save_interval=-1,
|
|
lambda_preserve=1.0, attn_chunk_size=4096, max_train_steps=-1,
|
|
skip_save_eval=False, repeat_single_sample=1, max_len=50000,
|
|
no_save=False, gradient_checkpointing=False, logits_to_keep=1):
|
|
# -------------------------------------------------------------------------
|
|
# GQA analysis
|
|
# -------------------------------------------------------------------------
|
|
num_q_heads = model.config.num_attention_heads
|
|
num_kv_heads = getattr(model.config, "num_key_value_heads", num_q_heads)
|
|
head_dim = model.config.hidden_size // num_q_heads
|
|
group_size = num_q_heads // num_kv_heads
|
|
|
|
print(f"📐 GQA config: {num_q_heads} Q heads, {num_kv_heads} KV heads, group_size={group_size}")
|
|
|
|
affected_structure = compute_gqa_affected_heads(target_structure, num_q_heads, num_kv_heads)
|
|
has_affected = len(affected_structure) > 0
|
|
if has_affected:
|
|
total_affected = sum(len(v) for v in affected_structure.values())
|
|
print(f"🛡️ Preservation targets: {total_affected} KV-group-affected unselected heads across {len(affected_structure)} layers (λ={lambda_preserve})")
|
|
for l, heads in sorted(affected_structure.items()):
|
|
print(f" Layer {l}: affected heads {heads} (selected: {target_structure[l]})")
|
|
else:
|
|
print("✅ No KV-group collateral — all touched KV groups are fully selected. No preservation loss needed.")
|
|
|
|
# -------------------------------------------------------------------------
|
|
# Setup LoRA
|
|
# -------------------------------------------------------------------------
|
|
layers = sorted(target_structure.keys())
|
|
targets = get_lora_targets(model, layers)
|
|
|
|
lora_cfg = LoraConfig(r=32, lora_alpha=16, bias="none", target_modules=targets, task_type="CAUSAL_LM")
|
|
gc_kwargs = {"use_reentrant": False} if gradient_checkpointing else None
|
|
model = prepare_model_for_kbit_training(
|
|
model,
|
|
use_gradient_checkpointing=gradient_checkpointing,
|
|
gradient_checkpointing_kwargs=gc_kwargs,
|
|
)
|
|
if gradient_checkpointing:
|
|
model.config.use_cache = False
|
|
print("🧩 Gradient checkpointing enabled (use_reentrant=False).")
|
|
|
|
if resume_adapter:
|
|
model = PeftModel.from_pretrained(model, resume_adapter, is_trainable=True)
|
|
print(f"♻️ Resumed LoRA: {resume_adapter}")
|
|
else:
|
|
model = get_peft_model(model, lora_cfg)
|
|
|
|
# -------------------------------------------------------------------------
|
|
# Apply head mask to LoRA layers (physical isolation)
|
|
# -------------------------------------------------------------------------
|
|
print("🎭 Applying GQA-aware head masks to LoRA layers...")
|
|
apply_head_mask_to_lora(model, target_structure, num_q_heads, num_kv_heads, head_dim)
|
|
|
|
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)")
|
|
|
|
# -------------------------------------------------------------------------
|
|
# Dataloader & Optimizer
|
|
# -------------------------------------------------------------------------
|
|
dl = DataLoader(
|
|
JsonlMessagesDS(data_path, repeat_single_sample=repeat_single_sample),
|
|
batch_size=bs,
|
|
shuffle=True,
|
|
collate_fn=lambda b: collate(b, tok),
|
|
)
|
|
opt = torch.optim.AdamW(model.parameters(), lr=lr)
|
|
sch = get_linear_schedule_with_warmup(opt, int(0.05 * epochs * len(dl)), epochs * len(dl))
|
|
|
|
current_idx = start_epoch
|
|
train_steps = 0
|
|
|
|
for ep in range(epochs):
|
|
if ep < start_epoch:
|
|
continue
|
|
pbar = tqdm(dl, desc=f"Ep {ep+1}/{epochs}")
|
|
for idx, batch in enumerate(pbar):
|
|
if ep == start_epoch and idx < start_batch_idx:
|
|
continue
|
|
device = next(model.parameters()).device
|
|
input_ids = batch["input_ids"][:, :max_len].to(device)
|
|
attn_mask = batch["attention_mask"][:, :max_len].to(device)
|
|
data_mask = batch["data_mask"][:, :max_len].to(device)
|
|
|
|
blind_attn_mask = attn_mask * (~data_mask).long()
|
|
|
|
# --- A. Base Model Pass (no LoRA, blind mask) ---
|
|
# Teacher for selected heads
|
|
model.eval()
|
|
base_map_container = {l: {h: None for h in h_list} for l, h_list in target_structure.items()}
|
|
model.config.retrieve_attn_map = base_map_container
|
|
model.config.retrieve_attn_chunk_size = attn_chunk_size
|
|
|
|
with torch.no_grad(), model.disable_adapter():
|
|
forward_for_attention_loss(model, input_ids, blind_attn_mask, logits_to_keep=logits_to_keep)
|
|
base_attns_map = base_map_container
|
|
del model.config.retrieve_attn_map
|
|
|
|
# --- A2. Original Model Pass (no LoRA, full mask) ---
|
|
# Teacher for KV-group-affected unselected heads
|
|
if has_affected:
|
|
orig_map_container = {l: {h: None for h in h_list} for l, h_list in affected_structure.items()}
|
|
model.config.retrieve_attn_map = orig_map_container
|
|
model.config.retrieve_attn_chunk_size = attn_chunk_size
|
|
|
|
with torch.no_grad(), model.disable_adapter():
|
|
forward_for_attention_loss(model, input_ids, attn_mask, logits_to_keep=logits_to_keep)
|
|
orig_attns_map = orig_map_container
|
|
del model.config.retrieve_attn_map
|
|
else:
|
|
orig_attns_map = None
|
|
|
|
# --- B. Tuned Model Pass (with LoRA, full mask) ---
|
|
# Accumulate chunked attention losses without materializing full tuned maps.
|
|
model.train()
|
|
loss_state = {
|
|
"selected": target_structure,
|
|
"affected": affected_structure,
|
|
"base_map": base_attns_map,
|
|
"orig_map": orig_attns_map,
|
|
"data_mask": data_mask,
|
|
"lambda_data": 1.0,
|
|
"eps": 1e-8,
|
|
"selected_losses": [],
|
|
"preserve_losses": [],
|
|
}
|
|
model.config.retrieve_attn_loss_state = loss_state
|
|
model.config.retrieve_attn_chunk_size = attn_chunk_size
|
|
|
|
forward_for_attention_loss(model, input_ids, attn_mask, logits_to_keep=logits_to_keep)
|
|
|
|
# --- C. Loss & Step ---
|
|
if loss_state["selected_losses"]:
|
|
loss_selected = torch.stack(loss_state["selected_losses"]).mean()
|
|
else:
|
|
loss_selected = torch.tensor(0.0, device=device)
|
|
|
|
loss_preserve = torch.tensor(0.0, device=device)
|
|
if has_affected and loss_state["preserve_losses"]:
|
|
loss_preserve = torch.stack(loss_state["preserve_losses"]).mean()
|
|
|
|
loss = loss_selected + lambda_preserve * loss_preserve
|
|
|
|
loss.backward()
|
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
|
opt.step()
|
|
sch.step()
|
|
opt.zero_grad()
|
|
|
|
del model.config.retrieve_attn_loss_state
|
|
pbar.set_postfix(
|
|
loss=f"{loss.item():.4f}",
|
|
sel=f"{loss_selected.item():.4f}",
|
|
pres=f"{loss_preserve.item():.4f}" if has_affected else "N/A",
|
|
)
|
|
if batch_save_interval > 0 and (idx + 1) % batch_save_interval == 0:
|
|
save_model(model, tok, out_dir, current_idx, idx, data_path, skip_eval=skip_save_eval)
|
|
|
|
train_steps += 1
|
|
if max_train_steps > 0 and train_steps >= max_train_steps:
|
|
if no_save:
|
|
print("⏭️ Skipped checkpoint save.")
|
|
else:
|
|
save_model(model, tok, out_dir, current_idx, idx, data_path, skip_eval=skip_save_eval)
|
|
print(f"✅ Reached max_train_steps={max_train_steps}.")
|
|
return
|
|
|
|
if no_save:
|
|
print("⏭️ Skipped checkpoint save.")
|
|
else:
|
|
save_model(model, tok, out_dir, current_idx, idx, data_path, skip_eval=skip_save_eval)
|
|
current_idx += 1
|
|
|
|
# -----------------------------------------------------------------------------
|
|
# 8. Main
|
|
# -----------------------------------------------------------------------------
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--model_path", default="/data/local/hujk/models/Llama-3.1-8B-Instruct")
|
|
parser.add_argument("--data_path", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-1_model_training_data_gen/single_turn/tri_native_tool_response_only.json")
|
|
parser.add_argument("--head_path", default="/home/hujk/gitrs/Paper2026/SortedCode2/2-2_head_identification_scoring/model_score/sep_Llama-3.1-8B-Instruct/heads_sorted/all_roc_inst_0.1.json")
|
|
|
|
parser.add_argument("--output_dir", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-2_model_training/outputs_lora")
|
|
parser.add_argument("--lora_path", default=None)
|
|
parser.add_argument("--epochs", type=int, default=3)
|
|
parser.add_argument("--batch_size", type=int, default=4)
|
|
parser.add_argument("--batch-save-interval", type=int, default=-1)
|
|
parser.add_argument("--lr", type=float, default=1e-4)
|
|
parser.add_argument("--topk", type=str, default="18.75p")
|
|
parser.add_argument("--lambda_preserve", type=float, default=1.0,
|
|
help="Weight for KV-group-affected preservation loss")
|
|
parser.add_argument("--attn-chunk-size", type=int, default=4096,
|
|
help="Sequence chunk size for last-token attention loss computation")
|
|
parser.add_argument("--max-train-steps", type=int, default=-1,
|
|
help="Stop after this many optimizer steps; <=0 means full training")
|
|
parser.add_argument("--skip-save-eval", action="store_true",
|
|
help="Save checkpoints without running MMLU/ASR evaluation")
|
|
parser.add_argument("--repeat-single-sample", type=int, default=1,
|
|
help="Repeat a one-record dataset in memory so batch-size tests can use real batches")
|
|
parser.add_argument("--max-len", type=int, default=50000,
|
|
help="Maximum token length kept from each batch")
|
|
parser.add_argument("--no-save", action="store_true",
|
|
help="Do not write checkpoints; useful for capacity smoke tests")
|
|
parser.add_argument("--gradient-checkpointing", action="store_true",
|
|
help="Enable transformer activation checkpointing for the trainable forward")
|
|
parser.add_argument("--logits-to-keep", type=int, default=1,
|
|
help="Forward only the last N logits; this loss does not need full-sequence logits")
|
|
args = parser.parse_args()
|
|
|
|
resume_adapter, start_epoch, start_batch_idx = discover_existing_adapter(args.output_dir)
|
|
if args.lora_path:
|
|
resume_adapter, start_epoch, start_batch_idx = args.lora_path, 0, 0
|
|
|
|
model, tok, tok_inf = load_model(args.model_path)
|
|
target_structure = get_target_structure(args.head_path, args.topk)
|
|
os.makedirs(args.output_dir, exist_ok=True)
|
|
tune_new(
|
|
model, tok, tok_inf, target_structure,
|
|
args.data_path, args.output_dir,
|
|
args.epochs, args.batch_size, args.lr,
|
|
resume_adapter=resume_adapter, start_epoch=start_epoch, start_batch_idx=start_batch_idx,
|
|
batch_save_interval=args.batch_save_interval,
|
|
lambda_preserve=args.lambda_preserve,
|
|
attn_chunk_size=args.attn_chunk_size,
|
|
max_train_steps=args.max_train_steps,
|
|
skip_save_eval=args.skip_save_eval,
|
|
repeat_single_sample=args.repeat_single_sample,
|
|
max_len=args.max_len,
|
|
no_save=args.no_save,
|
|
gradient_checkpointing=args.gradient_checkpointing,
|
|
logits_to_keep=args.logits_to_keep,
|
|
)
|
|
|
|
if __name__ == "__main__":
|
|
main()
|