karaka-attention / eval_baseline_jsd.py
rb512's picture
Upload folder using huggingface_hub
0de8760 verified
Raw
History Blame Contribute Delete
5.13 kB
# Baseline JSD: Standard MHA paraphrase stability (no Karaka layers)
# Compares attention head stability on paraphrase pairs using the raw encoder heads
import os, json, subprocess
import numpy as np
import torch
import torch.nn.functional as F
from transformers import AutoModel, AutoTokenizer
from peft import get_peft_model, LoraConfig, TaskType
from datasets import load_dataset
from tqdm.auto import tqdm
import torch_xla.core.xla_model as xm
DEVICE = xm.xla_device()
CONFIG = {
'base_model': 'Qwen/Qwen3-1.7B',
'lora_r': 16, 'lora_alpha': 32, 'lora_dropout': 0.05,
'lora_targets': ['q_proj', 'k_proj', 'v_proj', 'o_proj'],
'dhvani_checkpoint': 'gs://dhvani-checkpoints-vyasa/qwen3_lora_full_dhvani/dhvani_best.pt',
'max_length': 128,
'cache_dir': '/tmp/hf_cache',
'gcs_bucket': 'gs://karaka-checkpoints-vyasa/karaka_v1',
}
os.makedirs(CONFIG['cache_dir'], exist_ok=True)
def load_ckpt(path):
if path.startswith('gs://'):
local = f'/tmp/{os.path.basename(path)}'
if not os.path.exists(local):
subprocess.run(['gcloud', 'storage', 'cp', path, local], check=True)
return local
return path
@torch.no_grad()
def get_attention_maps(model, input_ids, attention_mask):
"""Get attention weights from the last transformer layer (standard MHA heads)."""
outputs = model(input_ids=input_ids, attention_mask=attention_mask, output_attentions=True)
# Last layer attention: [batch, num_heads, seq_len, seq_len]
last_attn = outputs.attentions[-1]
xm.mark_step()
return last_attn
def compute_jsd_per_head(attn1, attn2, num_heads=6):
"""Compute JSD between corresponding heads on two inputs."""
# attn shape: [1, num_heads, n, n] - use first 6 heads for fair comparison
jsds = []
for h in range(min(num_heads, attn1.shape[1])):
a1 = attn1[0, h].mean(0) + 1e-10 # avg over queries -> [n]
a2 = attn2[0, h].mean(0) + 1e-10
a1 = a1 / a1.sum()
a2 = a2 / a2.sum()
m = 0.5 * (a1 + a2)
jsd = 0.5 * (F.kl_div(m.log(), a1, reduction='sum', log_target=False) +
F.kl_div(m.log(), a2, reduction='sum', log_target=False))
jsds.append(max(0.0, jsd.item()))
return jsds
def main():
tokenizer = AutoTokenizer.from_pretrained(CONFIG['base_model'], trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# Load base model with LoRA (same as Dhvani encoder)
print('Loading model...')
base = AutoModel.from_pretrained(CONFIG['base_model'], torch_dtype=torch.bfloat16,
attn_implementation='eager', trust_remote_code=True)
lora_config = LoraConfig(r=CONFIG['lora_r'], lora_alpha=CONFIG['lora_alpha'],
lora_dropout=CONFIG['lora_dropout'],
target_modules=CONFIG['lora_targets'],
bias='none', task_type=TaskType.FEATURE_EXTRACTION)
model = get_peft_model(base, lora_config)
# Load Dhvani LoRA weights
print('Loading Dhvani checkpoint...')
ckpt = torch.load(load_ckpt(CONFIG['dhvani_checkpoint']), map_location='cpu')
model.load_state_dict(ckpt['lora'], strict=False)
del ckpt
model = model.to(DEVICE).eval()
# Load MRPC paraphrase pairs
ds = load_dataset('glue', 'mrpc', split='test', cache_dir=CONFIG['cache_dir'])
pairs = [(x['sentence1'], x['sentence2']) for x in ds if x['label'] == 1][:200]
print(f'Loaded {len(pairs)} paraphrase pairs')
# Compute JSD per head
all_jsds = [] # list of [6 heads] per pair
for s1, s2 in tqdm(pairs, desc='Baseline JSD'):
enc1 = tokenizer(s1, max_length=CONFIG['max_length'], truncation=True,
padding='max_length', return_tensors='pt')
enc2 = tokenizer(s2, max_length=CONFIG['max_length'], truncation=True,
padding='max_length', return_tensors='pt')
attn1 = get_attention_maps(model, enc1['input_ids'].to(DEVICE), enc1['attention_mask'].to(DEVICE))
attn2 = get_attention_maps(model, enc2['input_ids'].to(DEVICE), enc2['attention_mask'].to(DEVICE))
jsds = compute_jsd_per_head(attn1, attn2, num_heads=6)
all_jsds.append(jsds)
all_jsds = np.array(all_jsds) # [200, 6]
mean_per_head = all_jsds.mean(axis=0).tolist()
overall_mean = float(all_jsds.mean())
print(f'\nBaseline MHA JSD per head (first 6 heads of last layer):')
for i, jsd in enumerate(mean_per_head):
print(f' Head {i}: {jsd:.6f}')
print(f' Overall: {overall_mean:.6f}')
results = {
'baseline_mha_jsd_per_head': mean_per_head,
'baseline_mha_jsd_mean': overall_mean,
'n_pairs': len(pairs),
'note': 'Standard Qwen3-1.7B + Dhvani LoRA, last layer, first 6 heads',
}
out = '/tmp/baseline_jsd.json'
with open(out, 'w') as f:
json.dump(results, f, indent=2)
subprocess.run(['gcloud', 'storage', 'cp', out, f'{CONFIG["gcs_bucket"]}/baseline_jsd.json'])
print(f'Saved to GCS')
if __name__ == '__main__':
main()