Compress + absorb for arbitrary head groupings (e.g. similarity-driven). For each new GQLA head ``h_new = g*gs + i`` with original head ``orig_h = groups[g][i]``: slot_k = u_k[g, i*d_k:(i+1)*d_k, :] # (d_k, d_k) slot_v = u_v[g, i*d_v:(i+1)*d_v, :]
(
attn, layout: GqlaLayout, groups: list[list[int]],
u_k: torch.Tensor, u_v: torch.Tensor,
)
| 512 | |
| 513 | @torch.no_grad() |
| 514 | def compose_compressed_with_perm( |
| 515 | attn, layout: GqlaLayout, groups: list[list[int]], |
| 516 | u_k: torch.Tensor, u_v: torch.Tensor, |
| 517 | ) -> dict[str, torch.Tensor]: |
| 518 | """Compress + absorb for arbitrary head groupings (e.g. similarity-driven). |
| 519 | |
| 520 | For each new GQLA head ``h_new = g*gs + i`` with original head |
| 521 | ``orig_h = groups[g][i]``: |
| 522 | |
| 523 | slot_k = u_k[g, i*d_k:(i+1)*d_k, :] # (d_k, d_k) |
| 524 | slot_v = u_v[g, i*d_v:(i+1)*d_v, :] # (d_v, d_v) |
| 525 | new_W_q[h_new, nope] = slot_k.T @ W_q[orig_h, nope] |
| 526 | new_W_q[h_new, rope] = W_q[orig_h, rope] # rope just relabelled |
| 527 | new_W_o[:, h_new] = W_o[:, orig_h] @ slot_v |
| 528 | new_kv_b_K[g] = sum_i slot_k.T @ W_k[orig_h] |
| 529 | new_kv_b_V[g] = sum_i slot_v.T @ W_v[orig_h] |
| 530 | |
| 531 | Reduces to ``compress_and_absorb`` when ``groups = [[g*gs, g*gs+1, ...]]`` |
| 532 | (neighbor grouping) up to a (negligible) different reduction order. |
| 533 | """ |
| 534 | H, G, gs = layout.num_heads, layout.num_kv_heads, layout.group_size |
| 535 | d_k, d_rope, d_v, kv_lora = layout.qk_nope, layout.qk_rope, layout.v_dim, layout.kv_lora |
| 536 | qk_head_dim = d_k + d_rope |
| 537 | per_head_kv = d_k + d_v |
| 538 | |
| 539 | dev = attn.kv_b_proj.weight.device |
| 540 | u_k_d = u_k.to(device=dev, dtype=torch.float32) |
| 541 | u_v_d = u_v.to(device=dev, dtype=torch.float32) |
| 542 | |
| 543 | w_kv = attn.kv_b_proj.weight.data.to(torch.float32) # (H * per_head_kv, kv_lora) |
| 544 | w_per_head = w_kv.view(H, per_head_kv, kv_lora) |
| 545 | w_k_per_head = w_per_head[:, :d_k, :] # (H, d_k, kv_lora) |
| 546 | w_v_per_head = w_per_head[:, d_k:, :] # (H, d_v, kv_lora) |
| 547 | |
| 548 | new_kv_b = torch.zeros(G, per_head_kv, kv_lora, dtype=torch.float32, device=dev) |
| 549 | for g, grp in enumerate(groups): |
| 550 | acc_k = torch.zeros(d_k, kv_lora, dtype=torch.float32, device=dev) |
| 551 | acc_v = torch.zeros(d_v, kv_lora, dtype=torch.float32, device=dev) |
| 552 | for i, orig_h in enumerate(grp): |
| 553 | slot_k = u_k_d[g, i * d_k:(i + 1) * d_k, :] |
| 554 | slot_v = u_v_d[g, i * d_v:(i + 1) * d_v, :] |
| 555 | acc_k = acc_k + slot_k.T @ w_k_per_head[orig_h] |
| 556 | acc_v = acc_v + slot_v.T @ w_v_per_head[orig_h] |
| 557 | new_kv_b[g, :d_k, :] = acc_k |
| 558 | new_kv_b[g, d_k:, :] = acc_v |
| 559 | new_kv_b = new_kv_b.reshape(G * per_head_kv, kv_lora).contiguous() |
| 560 | |
| 561 | w_q = attn.q_b_proj.weight.data.to(torch.float32) # (H * qk_head_dim, q_lora) |
| 562 | new_w_q = torch.empty_like(w_q) |
| 563 | for g, grp in enumerate(groups): |
| 564 | for i, orig_h in enumerate(grp): |
| 565 | h_new = g * gs + i |
| 566 | slot_k = u_k_d[g, i * d_k:(i + 1) * d_k, :] |
| 567 | s_old = orig_h * qk_head_dim |
| 568 | s_new = h_new * qk_head_dim |
| 569 | new_w_q[s_new:s_new + d_k, :] = slot_k.T @ w_q[s_old:s_old + d_k, :] |
| 570 | new_w_q[s_new + d_k:s_new + qk_head_dim, :] = w_q[s_old + d_k:s_old + qk_head_dim, :] |
| 571 |