"""B3a-pivot: load ONE attention layer of Kimi K2.6 with REAL weights, run a single full-prefix forward over 256 synthetic tokens, capture the real attention distribution, then demonstrate the H2O eviction policy operating on it. This validates that the eviction policy makes sensible decisions when fed real Kimi K2.6 attention distributions (vs. random distributions in B1). Output: /tmp/kimi_layer_eviction_demo.csv with per-token attention scores + the eviction decision per token. """ import csv import json import sys import time from pathlib import Path import torch from safetensors import safe_open from transformers import DeepseekV3Config from transformers.models.deepseek_v3.modeling_deepseek_v3 import ( DeepseekV3Attention, DeepseekV3RotaryEmbedding, ) KIMI_PATH = Path("/mnt/llm_bank/Kimi-K2.6") SHARD = KIMI_PATH / "model-00001-of-000064.safetensors" LAYER_IDX = 0 LAYER_PREFIX = f"language_model.model.layers.{LAYER_IDX}.self_attn" SEQ_LEN = 256 BUDGET = 64 N_SINK = 4 N_RECENT = 32 def main() -> None: device = "cuda" if torch.cuda.is_available() else "cpu" print(f"[demo] device: {device}") if torch.cuda.is_available(): print(f"[demo] {torch.cuda.get_device_name(0)}") print(f"\n[demo] loading Kimi K2.6 config") full_config = json.load(open(KIMI_PATH / "config.json")) text_cfg = full_config["text_config"] cfg = DeepseekV3Config( vocab_size=text_cfg["vocab_size"], hidden_size=text_cfg["hidden_size"], intermediate_size=text_cfg["intermediate_size"], num_hidden_layers=text_cfg["num_hidden_layers"], num_attention_heads=text_cfg["num_attention_heads"], num_key_value_heads=text_cfg.get("num_key_value_heads", text_cfg["num_attention_heads"]), kv_lora_rank=text_cfg["kv_lora_rank"], q_lora_rank=text_cfg.get("q_lora_rank", 0) or 1536, qk_rope_head_dim=text_cfg["qk_rope_head_dim"], qk_nope_head_dim=text_cfg["qk_nope_head_dim"], v_head_dim=text_cfg["v_head_dim"], max_position_embeddings=text_cfg.get("max_position_embeddings", 4096), rope_theta=text_cfg.get("rope_theta", 10000.0), attn_implementation="eager", torch_dtype=torch.bfloat16, ) print(f"[demo] config: {cfg.num_hidden_layers}L hidden={cfg.hidden_size} heads={cfg.num_attention_heads}") print(f"[demo] qk_dim={cfg.qk_nope_head_dim+cfg.qk_rope_head_dim} v_dim={cfg.v_head_dim} kv_lora_rank={cfg.kv_lora_rank}") layer = DeepseekV3Attention(cfg, layer_idx=LAYER_IDX).to(dtype=torch.bfloat16) print(f"[demo] layer params: {sum(p.numel() for p in layer.parameters()):,}") print(f"\n[demo] loading layer-0 weights from {SHARD.name}") t0 = time.time() loaded = {} with safe_open(SHARD, framework="pt", device="cpu") as f: target_keys = [k for k in f.keys() if k.startswith(LAYER_PREFIX)] for k in target_keys: local_name = k[len(LAYER_PREFIX) + 1:] loaded[local_name] = f.get_tensor(k).to(dtype=torch.bfloat16) print(f"[demo] read {len(loaded)} weights in {time.time()-t0:.1f}s") missing, unexpected = layer.load_state_dict(loaded, strict=False) print(f"[demo] load_state_dict: missing={len(missing)} unexpected={len(unexpected)}") if missing: print(f" WARNING missing: {missing[:5]}") if unexpected: print(f" WARNING unexpected: {unexpected[:5]}") layer = layer.to(device).eval() rope = DeepseekV3RotaryEmbedding(config=cfg).to(device) # ---- Single full-prefix forward over SEQ_LEN tokens ---- print(f"\n[demo] running single forward over seq_len={SEQ_LEN}") bsz = 1 h = torch.randn(bsz, SEQ_LEN, cfg.hidden_size, dtype=torch.bfloat16, device=device) pos_ids = torch.arange(SEQ_LEN, dtype=torch.long, device=device).unsqueeze(0) # (1, SEQ_LEN) cos, sin = rope(h, pos_ids) # Causal attention mask: token i can attend to positions 0..i (lower-triangular). causal = torch.ones(SEQ_LEN, SEQ_LEN, dtype=torch.bool, device=device).tril() # Convert to additive mask: 0 where attend, -inf where masked attn_mask = torch.where( causal, torch.tensor(0.0, dtype=torch.bfloat16, device=device), torch.tensor(float("-inf"), dtype=torch.bfloat16, device=device), ) # Reshape to (bsz, 1, q_len, kv_len) attn_mask = attn_mask.unsqueeze(0).unsqueeze(0) t0 = time.time() try: with torch.no_grad(): out = layer( hidden_states=h, position_embeddings=(cos, sin), attention_mask=attn_mask, output_attentions=True, ) attn_out = out[0] if isinstance(out, tuple) else out attn_w = out[1] if isinstance(out, tuple) and len(out) > 1 else None except Exception as e: import traceback traceback.print_exc() sys.exit(1) print(f"[demo] forward done in {time.time()-t0:.2f}s") print(f"[demo] attn_out shape: {attn_out.shape}") print(f"[demo] attn_weights shape: {attn_w.shape if attn_w is not None else None}") if attn_w is None: print("[demo] no attention weights returned; cannot demonstrate eviction policy") sys.exit(1) # ---- Compute per-token cumulative attention mass (heavy-hitter score) ---- # attn_w shape: (bsz, num_heads, q_len, kv_len) # For each kv-position k, sum across heads and across all q that attend to k. # In a causal attention, position k receives attention from queries q >= k. score_per_token = attn_w[0].float().sum(dim=(0, 1)) # (kv_len,) score_per_token = score_per_token.cpu().numpy() print(f"[demo] score per token: shape={score_per_token.shape}") print(f" score range: [{score_per_token.min():.3f}, {score_per_token.max():.3f}]") print(f" score mean: {score_per_token.mean():.3f}") print(f" score std: {score_per_token.std():.3f}") # ---- Apply H2O eviction policy on REAL Kimi attention scores ---- print(f"\n[demo] applying H2O eviction: budget={BUDGET}, n_sink={N_SINK}, n_recent={N_RECENT}") sink_idx = list(range(N_SINK)) recent_idx = list(range(SEQ_LEN - N_RECENT, SEQ_LEN)) middle_range = list(range(N_SINK, SEQ_LEN - N_RECENT)) mid_with_score = [(i, float(score_per_token[i])) for i in middle_range] mid_with_score.sort(key=lambda x: -x[1]) heavy_idx = [i for i, _ in mid_with_score[:BUDGET]] keep = sorted(set(sink_idx) | set(heavy_idx) | set(recent_idx)) evict = [i for i in range(SEQ_LEN) if i not in keep] print(f"[demo] kept {len(keep)} of {SEQ_LEN} tokens ({100*len(keep)/SEQ_LEN:.1f}%)") print(f" sinks: {len(sink_idx)} (indices 0..{N_SINK-1})") print(f" heavy: {len(heavy_idx)} of {len(middle_range)} middle tokens chosen") print(f" recent: {len(recent_idx)} (last {N_RECENT})") print(f" evicted: {len(evict)} ({100*len(evict)/SEQ_LEN:.1f}%)") print(f" top 10 heavy-hitter scores: {[round(score_per_token[i], 2) for i in heavy_idx[:10]]}") # Sanity: heavy-hitters should have higher scores than the average evicted token if evict: evicted_mean = sum(score_per_token[i] for i in evict) / len(evict) kept_heavy_mean = sum(score_per_token[i] for i in heavy_idx) / len(heavy_idx) if heavy_idx else 0 print(f" mean score of heavy-hitters kept: {kept_heavy_mean:.3f}") print(f" mean score of evicted tokens: {evicted_mean:.3f}") ratio = kept_heavy_mean / max(evicted_mean, 1e-9) print(f" heavy/evicted score ratio: {ratio:.2f}x") # ---- Save per-token CSV for downstream analysis / plot ---- out_path = Path("/tmp/kimi_layer_eviction_demo.csv") with open(out_path, "w", newline="") as f: writer = csv.DictWriter(f, fieldnames=["token_idx", "attention_score", "kept", "category"]) writer.writeheader() for i in range(SEQ_LEN): if i in sink_idx: cat = "sink" elif i in recent_idx: cat = "recent" elif i in heavy_idx: cat = "heavy" else: cat = "evicted" writer.writerow({ "token_idx": i, "attention_score": round(float(score_per_token[i]), 6), "kept": cat != "evicted", "category": cat, }) print(f"\n[demo] wrote {SEQ_LEN} rows -> {out_path}") print("[demo] DONE") if __name__ == "__main__": main()