"""Optimized local inference: Fuse-2 QLoRA + DSpark speculative decoding. Key optimizations over the deepspec default loop: 1. Forward hooks on 5 layers instead of output_hidden_states=True (all 36) 2. Greedy argmax verification (no rejection sampling, no strict fallback) 3. Single target forward pass per block (no sequential fallback) 4. Efficient KV cache management This makes DSpark provide speedup on ANY GPU + quantization combo. Usage: D:\\fuse2-venv\\Scripts\\python.exe fuse2_dspark_fast.py "your prompt" 512 """ import os import sys import gc import json import time from types import SimpleNamespace # Use expandable segments to reduce CUDA reserved memory overhead. # This can save ~0.5 GB of reserved VRAM, which is the difference between # fitting in 12 GB and overflowing to PCIe (50x slowdown). os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import torch import torch.nn as nn import torch.nn.functional as F sys.stdout.reconfigure(encoding='utf-8', errors='replace') sys.stderr.reconfigure(encoding='utf-8', errors='replace') sys.path.insert(0, "E:/fuse1/DeepSpec") sys.path.insert(0, "E:/fuse1/mini-deepseek-v4-flash-qlora") MODEL_PATH = "E:/fuse1/mini-deepseek-v4-flash-qlora" DRAFT_DIR = "E:/fuse1/mini-deepseek-v4-flash/drafter" DEVICE = torch.device("cuda:0") BLOCK_SIZE = 7 # Optimal for this drafter (trained with block_size=7) SWIGLU_LIMIT = 10.0 # ── Load target (4-bit NF4) ───────────────────────────────────────────────── def load_target(): from transformers import AutoConfig, AutoTokenizer from accelerate import init_empty_weights from safetensors import safe_open import bitsandbytes as bnb import fuse2_model_local print(f"[target] Loading {MODEL_PATH}...") config = AutoConfig.from_pretrained(MODEL_PATH, trust_remote_code=True) # Force SDPA (fused CUDA kernel) instead of eager (Python loop) config._attn_implementation = "sdpa" with init_empty_weights(): model = fuse2_model_local.Fuse2ForCausalLM(config) with open(f"{MODEL_PATH}/model.safetensors.index.json") as f: index = json.load(f) weight_map = index["weight_map"] shards = sorted(set(weight_map.values())) param_names = set(dict(model.named_parameters()).keys()) buffer_names = set(dict(model.named_buffers()).keys()) tok = AutoTokenizer.from_pretrained(MODEL_PATH, trust_remote_code=True) if tok.pad_token_id is None: tok.pad_token_id = tok.eos_token_id # Test: BF16 attention + shared drafter embedding to save VRAM # BF16 attention is 2x faster than 4-bit for small (2560x2560) matrices quant_suffixes = ("gate_proj.weight", "up_proj.weight", "down_proj.weight") def nav(model, key): parts = key.split(".") obj = model for p in parts[:-1]: obj = obj[int(p)] if p.isdigit() else getattr(obj, p) return obj, parts[-1] def find_linear(model, key): parts = key.split(".") obj = model for p in parts[:-2]: obj = obj[int(p)] if p.isdigit() else getattr(obj, p) return obj, parts[-2] replaced = {} for shard_name in shards: with safe_open(os.path.join(MODEL_PATH, shard_name), framework="pt", device="cpu") as f: for key in [k for k, v in weight_map.items() if v == shard_name]: if key not in param_names and key not in buffer_names: continue tensor = f.get_tensor(key) if any(key.endswith(s) for s in quant_suffixes) and key in param_names: owner, attr = find_linear(model, key) lid = id(getattr(owner, attr)) if lid not in replaced: old = getattr(owner, attr) new = bnb.nn.Linear4bit(old.in_features, old.out_features, bias=old.bias is not None, quant_type="nf4", compute_dtype=torch.bfloat16, device=DEVICE) t = tensor.to(torch.bfloat16) new.weight = bnb.nn.Params4bit(t, requires_grad=False, quant_type="nf4").cuda(0) if old.bias is not None: new.bias = None setattr(owner, attr, new) replaced[id(new)] = new del old, new, t else: parent, pname = nav(model, key) parent._parameters[pname] = nn.Parameter(tensor.to(torch.bfloat16).to(DEVICE), requires_grad=False) del tensor gc.collect() torch.cuda.empty_cache() if model.lm_head.weight.device.type == 'meta': model.lm_head.weight = nn.Parameter(model.model.embed_tokens.weight.data.clone(), requires_grad=False) from fuse2_model_local import Fuse2AugmentedLayer for layer in model.model.layers: if not isinstance(layer, Fuse2AugmentedLayer): continue if hasattr(layer, 'coding_gate') and layer.coding_gate.device.type == 'meta': layer.coding_gate = nn.Parameter(torch.tensor(-2.0, device=DEVICE)) if hasattr(layer, 'coding_norm') and layer.coding_norm.weight.device.type == 'meta': layer.coding_norm = nn.RMSNorm(layer.coding_norm.weight.shape[0], eps=1e-6).to(DEVICE) experts = getattr(layer, "experts", None) if experts: for expert in experts: gp, up, dp = expert.gate_proj, expert.up_proj, expert.down_proj def make_fwd(g, u, d, lim): def forward(x): return d(torch.clamp(F.silu(g(x)) * u(x), -lim, lim)) return forward # Try torch.compile for kernel fusion (silu+mul+clamp in one kernel) # Disabled: triton not available on Windows, torch.compile fails # try: # expert.forward = torch.compile(make_fwd(gp, up, dp, SWIGLU_LIMIT), mode="reduce-overhead") # except Exception: expert.forward = make_fwd(gp, up, dp, SWIGLU_LIMIT) router = getattr(layer, "router", None) if router: gate, top_k = router.gate, router.top_k def make_router(g, tk): def forward(h): logits = g(h) scores = torch.clamp(F.softplus(logits), min=1e-6).sqrt() w, idx = scores.topk(tk, dim=-1) return w / (w.sum(dim=-1, keepdim=True) + 1e-8), idx, logits return forward router.forward = make_router(gate, top_k) model.set_coding_enabled(True) model.to(DEVICE) model.eval() # Fix norm dtypes: convert all RMSNorm/LayerNorm weights to bfloat16 # (they load as float32, blocking fused kernel dispatch) norm_count = 0 for module in model.modules(): if hasattr(module, 'weight') and hasattr(module, 'eps') and module.weight.dtype == torch.float32: module.weight.data = module.weight.data.to(torch.bfloat16) norm_count += 1 print(f"[target] Fixed {norm_count} norm layers to BF16") # Verify attention implementation print(f"[target] Attention: {config._attn_implementation}") print(f"[target] VRAM: {torch.cuda.memory_allocated(0)/1e9:.1f} GB") return model, tok # ── Load drafter ──────────────────────────────────────────────────────────── def load_drafter(): from deepspec.modeling.dspark.qwen3 import Qwen3DSparkModel import bitsandbytes as bnb print(f"[draft] Loading {DRAFT_DIR}...") # Load in BF16 first, then quantize linear layers to 4-bit to save VRAM. # Drafter is 2.8 GB in BF16 → ~0.7 GB in 4-bit, saving 2.1 GB. # This keeps total VRAM at ~9.6 GB (well under 12 GB limit). draft = Qwen3DSparkModel.from_pretrained(DRAFT_DIR, dtype=torch.bfloat16, attn_implementation="sdpa").to(DEVICE).eval() # Quantize all Linear layers in drafter to 4-bit q_count = 0 for name, module in draft.named_modules(): if isinstance(module, torch.nn.Linear) and module.weight.device.type != 'meta': old_w = module.weight.data if old_w.dtype == torch.bfloat16: p4bit = bnb.nn.Params4bit(old_w.cpu(), requires_grad=False, quant_type="nf4").to(DEVICE) new_linear = bnb.nn.Linear4bit( module.in_features, module.out_features, bias=module.bias is not None, quant_type="nf4", compute_dtype=torch.bfloat16, device=DEVICE, ) new_linear.weight = p4bit if module.bias is not None: new_linear.bias = module.bias # Replace in parent parent_name = name.rsplit(".", 1)[0] if "." in name else "" attr_name = name.rsplit(".", 1)[-1] if "." in name else name parent = draft if parent_name: for p in parent_name.split("."): parent = getattr(parent, p) setattr(parent, attr_name, new_linear) q_count += 1 draft.block_size = BLOCK_SIZE target_layer_ids = [int(x) for x in draft.target_layer_ids] print(f"[draft] Quantized {q_count} layers to 4-bit") # Aggressive cleanup: the BF16→4-bit conversion leaves orphaned tensors gc.collect() torch.cuda.empty_cache() torch.cuda.synchronize() print(f"[draft] VRAM: {torch.cuda.memory_allocated(0)/1e9:.1f} GB " f"(reserved: {torch.cuda.memory_reserved(0)/1e9:.1f} GB), " f"layers: {target_layer_ids}") return draft, target_layer_ids def load_drafter_shared(target_embed): """Load drafter with shared embedding from target to save ~0.4 GB VRAM.""" from deepspec.modeling.dspark.qwen3 import Qwen3DSparkModel import bitsandbytes as bnb print(f"[draft] Loading {DRAFT_DIR} (shared embedding)...") draft = Qwen3DSparkModel.from_pretrained(DRAFT_DIR, dtype=torch.bfloat16, attn_implementation="sdpa").to(DEVICE).eval() # Share embedding with target (same vocab) if hasattr(draft, 'embed_tokens') and target_embed is not None: old_embed = draft.embed_tokens draft.embed_tokens = target_embed del old_embed gc.collect() torch.cuda.empty_cache() print(f"[draft] Shared embedding with target") # Quantize all Linear layers in drafter to 4-bit q_count = 0 for name, module in draft.named_modules(): if isinstance(module, torch.nn.Linear) and module.weight.device.type != 'meta': old_w = module.weight.data if old_w.dtype == torch.bfloat16: p4bit = bnb.nn.Params4bit(old_w.cpu(), requires_grad=False, quant_type="nf4").to(DEVICE) new_linear = bnb.nn.Linear4bit( module.in_features, module.out_features, bias=module.bias is not None, quant_type="nf4", compute_dtype=torch.bfloat16, device=DEVICE, ) new_linear.weight = p4bit if module.bias is not None: new_linear.bias = module.bias parent_name = name.rsplit(".", 1)[0] if "." in name else "" attr_name = name.rsplit(".", 1)[-1] if "." in name else name parent = draft if parent_name: for p in parent_name.split("."): parent = getattr(parent, p) setattr(parent, attr_name, new_linear) q_count += 1 draft.block_size = BLOCK_SIZE target_layer_ids = [int(x) for x in draft.target_layer_ids] print(f"[draft] Quantized {q_count} layers to 4-bit") gc.collect() torch.cuda.empty_cache() torch.cuda.synchronize() print(f"[draft] VRAM: {torch.cuda.memory_allocated(0)/1e9:.1f} GB " f"(reserved: {torch.cuda.memory_reserved(0)/1e9:.1f} GB), " f"layers: {target_layer_ids}") return draft, target_layer_ids # ── Hidden state capture via hooks (only 5 layers, not 36) ────────────────── class HiddenStateCapture: """Capture hidden states from specific layers using forward hooks. Much faster than output_hidden_states=True which stores all 36 layers.""" def __init__(self, model, layer_ids): self.layer_ids = layer_ids self.captured = {} self.hooks = [] for lid in layer_ids: layer = model.model.layers[lid] h = layer.register_forward_hook(self._make_hook(lid)) self.hooks.append(h) def _make_hook(self, lid): def hook(module, input, output): # output is a tuple (hidden_states, ...) or just hidden_states if isinstance(output, tuple): hs = output[0] else: hs = output self.captured[lid] = hs.detach() return hook def get_context_feature(self): """Return concatenated hidden states for the drafter. Matches extract_context_feature: cat([hs[lid] for lid in layer_ids], dim=-1) """ return torch.cat([self.captured[lid] for lid in self.layer_ids], dim=-1) def reset(self): self.captured.clear() def remove(self): for h in self.hooks: h.remove() # ── Optimized speculative decode loop ─────────────────────────────────────── def generate_fast(target, draft, tok, prompt, max_new_tokens=1024): """Custom speculative decode loop optimized for any GPU + quantization. Key differences from deepspec's generate_decoding_sample: 1. Forward hooks on 5 layers instead of output_hidden_states=True (all 36) 2. Greedy argmax verification (no rejection sampling overhead) 3. NO sequential fallback — single target forward per block 4. Minimal Python overhead per block """ from transformers import DynamicCache from deepspec.eval.dspark.draft_ops import ( build_dspark_proposal, forward_dspark_draft_block, ) draft_dtype = torch.bfloat16 device = DEVICE # Format prompt messages = [{"role": "user", "content": prompt}] text = tok.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) input_ids = tok(text, return_tensors="pt").input_ids.to(device) input_len = input_ids.shape[1] eos_id = int(tok.eos_token_id) # Global position_ids tensor (covers all possible positions) max_length = input_len + max_new_tokens + BLOCK_SIZE + 1 global_position_ids = torch.arange(max_length, device=device).unsqueeze(0) # Set up hidden state capture hooks on target capture = HiddenStateCapture(target, target_layer_ids) # ── Benchmark: standalone target forward (no drafter interference) ── print(f"[gen] Benchmark: standalone target forward...", flush=True) import gc as gc_mod gc_mod.collect() torch.cuda.empty_cache() print(f" VRAM allocated: {torch.cuda.memory_allocated(0)/1e9:.2f} GB, " f"reserved: {torch.cuda.memory_reserved(0)/1e9:.2f} GB", flush=True) bench_cache = DynamicCache() bench_ids = tok("Hello world", return_tensors="pt").input_ids.to(device) with torch.inference_mode(): # Warmup for _ in range(3): out = target(input_ids=bench_ids, past_key_values=bench_cache, use_cache=True) bench_cache.crop(bench_cache.get_seq_length() - bench_ids.shape[1]) torch.cuda.synchronize() t0 = time.perf_counter() test8 = torch.tensor([[out.logits[:, -1:].argmax(dim=-1).item()]*8], device=device) for _ in range(5): out = target(input_ids=test8, past_key_values=bench_cache, use_cache=True) bench_cache.crop(bench_cache.get_seq_length() - 8) torch.cuda.synchronize() bench_time = (time.perf_counter() - t0) / 5 * 1000 print(f" 8-token forward: {bench_time:.0f}ms ({8*5/(bench_time/1000*5):.1f} tok/s)", flush=True) del bench_cache, bench_ids, test8, out gc_mod.collect() torch.cuda.empty_cache() # ── Prefill ── print(f"[gen] Prefill ({input_len} tokens)...", flush=True) target_cache = DynamicCache() draft_cache = DynamicCache() with torch.inference_mode(): # Target prefill (hooks capture hidden states) capture.reset() prefill_output = target( input_ids=input_ids, position_ids=global_position_ids[:, :input_len], past_key_values=target_cache, use_cache=True, ) # First token (greedy) next_token = prefill_output.logits[:, -1:].argmax(dim=-1) # Extract hidden states for drafter target_hidden = capture.get_context_feature().to(dtype=draft_dtype) # Build output buffer output_tokens = [next_token.item()] start = input_len # ── Speculative decode loop ── block_count = 0 total_accepted = 0 fallback_count = 0 print(f"[gen] Speculative decoding (block={BLOCK_SIZE})...", flush=True) t0 = time.perf_counter() target_time = 0.0 draft_time = 0.0 while len(output_tokens) < max_new_tokens: block_count += 1 # ── 1. Drafter proposes BLOCK_SIZE tokens ── t_draft_start = time.perf_counter() draft_ids = torch.full((1, BLOCK_SIZE), int(draft.mask_token_id), dtype=torch.long, device=device) draft_ids[:, 0] = next_token block_hidden = forward_dspark_draft_block( draft, draft_input_ids=draft_ids, position_ids=global_position_ids, past_key_values_draft=draft_cache, target_hidden_states=target_hidden, start=start, block_size=BLOCK_SIZE, ) proposal = build_dspark_proposal( model=draft, draft_input_ids=draft_ids, block_hidden=block_hidden, block_size=BLOCK_SIZE, temperature=0.0, confidence_threshold=0.0, ) torch.cuda.synchronize() draft_time += time.perf_counter() - t_draft_start draft_count = int(proposal.draft_token_count) if draft_count == 0: # No draft tokens — just use target for next token capture.reset() out = target( input_ids=next_token.unsqueeze(0), position_ids=global_position_ids[:, start:start+1], past_key_values=target_cache, use_cache=True, ) next_token = out.logits[:, -1:].argmax(dim=-1) target_hidden = capture.get_context_feature().to(dtype=draft_dtype) output_tokens.append(next_token.item()) start += 1 if next_token.item() == eos_id: break continue # ── 2. Target verifies (ONE forward pass on all draft tokens) ── t_target_start = time.perf_counter() verify_ids = proposal.verify_input_ids # (1, draft_count + 1) verify_len = verify_ids.shape[1] t_hook_reset = time.perf_counter() capture.reset() torch.cuda.synchronize() t_reset_done = time.perf_counter() t_fwd_start = time.perf_counter() verify_output = target( input_ids=verify_ids, position_ids=global_position_ids[:, start:start + verify_len], past_key_values=target_cache, use_cache=True, ) torch.cuda.synchronize() t_fwd_done = time.perf_counter() t_extract_start = time.perf_counter() target_hidden = capture.get_context_feature().to(dtype=draft_dtype) torch.cuda.synchronize() t_extract_done = time.perf_counter() target_time += time.perf_counter() - t_target_start if block_count <= 5 or block_count % 20 == 0: print(f" target: reset={t_reset_done-t_hook_reset:.3f}s, " f"fwd={t_fwd_done-t_fwd_start:.3f}s, " f"extract={t_extract_done-t_extract_start:.3f}s, " f"vlen={verify_len}", flush=True) # ── 3. Greedy argmax verification ── # Target logits for positions 0..draft_count-1 predict tokens 1..draft_count target_logits = verify_output.logits[:, :draft_count, :] target_tokens = target_logits.argmax(dim=-1) # (1, draft_count) proposed_tokens = verify_ids[:, 1:1 + draft_count] # (1, draft_count) # Find first mismatch matches = (target_tokens == proposed_tokens).squeeze(0) # (draft_count,) # Accept prefix until first mismatch if matches.all(): accepted = draft_count else: # cumprod gives 1s until first 0, then 0s accept_mask = matches.cumprod(dim=0) accepted = int(accept_mask.sum().item()) total_accepted += accepted # ── 4. Commit accepted tokens + 1 bonus from target ── # The bonus token is the target's prediction at position `accepted` if accepted < draft_count: # Mismatch at position `accepted` — target's token there is the bonus bonus_token = target_tokens[:, accepted:accepted+1] else: # All accepted — bonus is target's prediction at the last position bonus_token = verify_output.logits[:, draft_count:draft_count+1, :].argmax(dim=-1) # Update KV cache: crop to start + accepted + 1 (we keep the bonus token's KV) new_start = start + accepted + 1 target_cache.crop(new_start) # Sliding window: limit KV cache to MAX_KV_CACHE tokens to prevent # speed degradation on long sequences. This crops the oldest entries # while keeping the most recent context. MAX_KV_CACHE = 512 # tokens cache_len = target_cache.get_seq_length() if cache_len > MAX_KV_CACHE: crop_to = cache_len - MAX_KV_CACHE for layer in target_cache.layers: if hasattr(layer, 'keys') and layer.keys.numel() > 0: layer.keys = layer.keys[:, :, crop_to:, :] layer.values = layer.values[:, :, crop_to:, :] # Update drafter hidden states from the captured hooks target_hidden = capture.get_context_feature().to(dtype=draft_dtype) # Only keep hidden states for accepted + 1 positions target_hidden = target_hidden[:, :accepted + 1, :] # Update draft cache to match accepted tokens draft_cache.crop(start + accepted) # Record tokens for i in range(accepted): output_tokens.append(int(proposed_tokens[0, i].item())) output_tokens.append(int(bonus_token[0, 0].item())) next_token = bonus_token start = new_start # Check for EOS if int(bonus_token[0, 0].item()) == eos_id: break if block_count % 20 == 0: elapsed = time.perf_counter() - t0 tps = len(output_tokens) / elapsed mean_acc = total_accepted / block_count avg_target = target_time / block_count * 1000 avg_draft = draft_time / block_count * 1000 print(f" [block {block_count}] {len(output_tokens)} tok, " f"{tps:.1f} tok/s, accept {mean_acc:.1f}/{BLOCK_SIZE}, " f"target {avg_target:.0f}ms, draft {avg_draft:.0f}ms", flush=True) elapsed = time.perf_counter() - t0 capture.remove() # Decode token_ids = torch.tensor(output_tokens, dtype=torch.long, device=device).unsqueeze(0) text_out = tok.decode(token_ids[0], skip_special_tokens=True) mean_accept = total_accepted / max(block_count, 1) tps = len(output_tokens) / max(elapsed, 1e-9) print(f"\n[gen] {len(output_tokens)} tokens in {elapsed:.1f}s ({tps:.1f} tok/s)") print(f"[gen] {block_count} blocks, mean accepted: {mean_accept:.2f}/{BLOCK_SIZE}") return text_out # ── Main ──────────────────────────────────────────────────────────────────── if __name__ == "__main__": prompt = sys.argv[1] if len(sys.argv) > 1 else "Write a Python fizzbuzz." max_tokens = int(sys.argv[2]) if len(sys.argv) > 2 else 2048 print("=" * 60) print("Fuse-2 QLoRA + DSpark Fast Speculative (Local)") print("=" * 60) target, tok = load_target() draft, target_layer_ids = load_drafter_shared(target.model.embed_tokens) vram = torch.cuda.memory_allocated(0) / 1e9 print(f"\n[ready] Total VRAM: {vram:.1f} GB") print("\n" + "=" * 60) print("GENERATING") print("=" * 60) result = generate_fast(target, draft, tok, prompt, max_new_tokens=max_tokens) print("\n" + "=" * 60) print("OUTPUT") print("=" * 60) print(result)