samai-8b-M8 / artifacts /r13_scripts /patch_refine_v2.py
tchbcb's picture
r13: slabbed refine_tensor2 (diag-H exact, VRAM-safe)
42846dd verified
Raw History Blame Contribute Delete
3.83 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""patch_refine_v2.py — refine_tensor 分 slab 显存安全版
对角 H → 损失逐块可分 → 16M 权重/slab 独立优化(数学等价), 峰值 VRAM ~3GB (原 ~13GB 驱动级 OOM)
"""
NEWFN = '''def refine_tensor2(W, hv, typ, steps, dev, wps=16000000):
"""W (out,in) fp32 cuda; hv (in,) normalized diag weights. Slab-chunked GSQ-lite."""
import torch
div = 8.0 if typ == "q4_0" else 16.0
out_ch, in_ch = W.shape
blk = 32
Wb = W.reshape(-1, blk)
nb = Wb.shape[0]
hvf = hv.to(dev).float().clamp_min(1e-12)
hvn = hvf / hvf.sum()
hv_blk = hvn.repeat(out_ch, 1).reshape(-1, blk)
codes_full = torch.empty(nb, blk, dtype=torch.int32, device=dev)
d_full = torch.empty(nb, dtype=torch.float32, device=dev)
bpb = max(1, (wps // 4) // blk)
n_slabs = (nb + bpb - 1) // bpb
for s in range(n_slabs):
b0, b1 = s * bpb, min(nb, (s + 1) * bpb)
Ws = Wb[b0:b1]
ns = Ws.shape[0]
hbs = hv_blk[b0:b1]
d0 = Ws.abs().amax(dim=1) / div
d0 = torch.where(d0 < 1e-12, torch.ones_like(d0), d0)
c0 = torch.clamp(torch.round(Ws / d0.unsqueeze(1)), -div, div - 1)
grid = torch.stack([c0 - 2.0 + k for k in range(5)], dim=-1)
grid = torch.clamp(grid, -div, div - 1)
logits = torch.zeros(ns, blk, 5, device=dev, requires_grad=True)
logits.data[:, :, 2] = 0.5
log_d = torch.log(d0).to(dev).requires_grad_(True)
opt = torch.optim.AdamW([
{"params": [logits], "lr": 3e-3},
{"params": [log_d], "lr": 1e-3, "weight_decay": 0.0}])
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=steps)
d0min, d0max = 0.3 * d0, 3.0 * d0
for ep in range(steps):
prog = ep / max(1, steps - 1)
tau = 2.0 * (0.05 / 2.0) ** prog
kap = 200.0 ** prog
d = torch.clamp(torch.exp(log_d), d0min, d0max)
g = -torch.log(-torch.log(torch.rand_like(logits) + 1e-9) + 1e-9)
ys = torch.softmax((kap * logits + g) / tau, dim=-1)
yh = torch.nn.functional.one_hot(ys.argmax(-1), 5).float()
y = yh.detach() - ys.detach() + ys
Q = (y * grid).sum(-1) * d.unsqueeze(1)
E = Ws - Q
hl = ((E * E) * hbs).sum()
mse = (E * E).mean()
loss = hl + 0.02 * mse
opt.zero_grad(set_to_none=True)
loss.backward()
if prog < 0.25:
log_d.grad = None
opt.step()
sched.step()
with torch.no_grad():
d = torch.clamp(torch.exp(log_d), d0min, d0max)
yh = torch.nn.functional.one_hot(logits.argmax(-1), 5).float()
codes_s = (yh * (grid + div)).sum(-1).round().clamp(0, 2 * div - 1).to(torch.int32)
codes_full[b0:b1] = codes_s
d_full[b0:b1] = d
del opt, sched, logits, log_d, grid, ys, yh, y, g, Q, E, loss
if dev == "cuda":
torch.cuda.empty_cache()
return d_full.cpu(), codes_full.reshape(-1).cpu(), int(steps)
'''
CALL_OLD = ''' d, codes, st = refine_tensor(W, H, typ, steps, dev)'''
CALL_NEW = ''' hv = torch.diagonal(H).clone().clamp_min(1e-12)
d, codes, st = refine_tensor2(W, hv, typ, steps, dev)'''
def main():
path = "/tmp/k8b/r13_refine.py"
src = open(path).read()
anchor = "def np_pack_blocks(d_np, codes_np, typ):"
assert src.count(anchor) == 1, "fn anchor %d" % src.count(anchor)
src = src.replace(anchor, NEWFN + anchor, 1)
assert src.count(CALL_OLD) == 1, "call anchor %d" % src.count(CALL_OLD)
src = src.replace(CALL_OLD, CALL_NEW, 1)
compile(src, path, "exec")
open(path, "w").write(src)
print("PATCH_OK refine_tensor2 slabbed")
main()