lmc-code / temp /gptj.py
khanhvinh9's picture
Upload folder using huggingface_hub
a20151e verified
Raw
History Blame Contribute Delete
2.92 kB
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}")