import torch import numpy as np import csv from transformers import GPTJModel # Load GPT-J (weights on CPU with float16) model = GPTJModel.from_pretrained( "EleutherAI/gpt-j-6b", torch_dtype=torch.float16, low_cpu_mem_usage=True ).cpu() hidden_size = model.config.hidden_size num_heads = model.config.n_head head_dim = hidden_size // num_heads # ---- Norm calculator ---- def compute_norms(A: torch.Tensor): A = A.float() # for numerical stability return { "1": torch.norm(A, p=1).item(), "F": torch.norm(A, p="fro").item(), "*": torch.linalg.svdvals(A).sum().item(), "2,1": torch.norm(A, dim=0, p=2).sum().item(), "2,1,T": torch.norm(A.t(), dim=0, p=2).sum().item() } norm_names = ["1","F","*","2,1","2,1,T"] subcols = ["Q","K","Q/K","V","O","V/O"] outfile = "gptj_qkvo_norms.csv" with open(outfile, "w", newline="") as f: writer = csv.writer(f) for layer_idx, layer in enumerate(model.h, start=1): # ---- Headers ---- header1 = [f"Layer {layer_idx}"] for n in norm_names: header1.extend([n,"","","","",""]) writer.writerow(header1) header2 = [""] for _ in norm_names: header2.extend(subcols) writer.writerow(header2) rows = [] # ---- Extract weights (no bias in GPT-J) ---- W_q = layer.attn.q_proj.weight.detach().cpu() W_k = layer.attn.k_proj.weight.detach().cpu() W_v = layer.attn.v_proj.weight.detach().cpu() W_o = layer.attn.out_proj.weight.detach().cpu() # ---- Split into heads ---- W_q_heads = W_q.view(num_heads, head_dim, -1) W_k_heads = W_k.view(num_heads, head_dim, -1) W_v_heads = W_v.view(num_heads, head_dim, -1) W_o_heads = W_o.view(num_heads, head_dim, -1) # ---- Per-head ---- for h in range(num_heads): row = [f"Head {h+1}"] for norm in norm_names: nq = compute_norms(W_q_heads[h])[norm] nk = compute_norms(W_k_heads[h])[norm] nv = compute_norms(W_v_heads[h])[norm] no = compute_norms(W_o_heads[h])[norm] qk_ratio = nq / (nk + 1e-12) vo_ratio = nv / (no + 1e-12) row.extend([ round(nq,4), round(nk,4), round(qk_ratio,4), round(nv,4), round(no,4), round(vo_ratio,4) ]) writer.writerow(row) rows.append(row[1:]) # ---- Mean & Std ---- arr = np.array(rows, dtype=float) mean = np.round(arr.mean(axis=0),4) std = np.round(arr.std(axis=0),4) writer.writerow(["Mean"] + mean.tolist()) writer.writerow(["Std"] + std.tolist()) writer.writerow([]) print(f"✅ Saved CSV: {outfile}")