MCPcopy Create free account
hub / github.com/MuLabPKU/TransArch / compose_compressed_with_perm

Function compose_compressed_with_perm

GQLA_preprint/src/compression.py:514–586  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

512
513@torch.no_grad()
514def 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

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected