"""Quick local inference for the QLoRA-merged Fuse-2 model. Loads the model in 4-bit NF4 (fits in 12 GB VRAM), applies the runtime patches (expert scale, SwiGLU clamp, router stability), and generates text for a given prompt. """ import os import sys import gc import json import time 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/mini-deepseek-v4-flash-qlora") MODEL_PATH = "E:/fuse1/mini-deepseek-v4-flash-qlora" DEVICE = "cuda:0" SWIGLU_LIMIT = 10.0 TARGET_STD = 0.025 torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.benchmark = True def load_model_4bit(): 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"[load] Config from {MODEL_PATH}...") config = AutoConfig.from_pretrained(MODEL_PATH, trust_remote_code=True) 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 quantize_suffixes = ( "q_proj.weight", "k_proj.weight", "v_proj.weight", "o_proj.weight", "gate_proj.weight", "up_proj.weight", "down_proj.weight", ) def should_quantize(key): return any(key.endswith(s) for s in quantize_suffixes) def navigate_to_parent(model, key): parts = key.split(".") obj = model for p in parts[:-1]: if p.isdigit(): obj = obj[int(p)] else: obj = getattr(obj, p) return obj, parts[-1] def find_linear_owner(model, key): parts = key.split(".") obj = model for p in parts[:-2]: if p.isdigit(): obj = obj[int(p)] else: obj = getattr(obj, p) return obj, parts[-2] print(f"[load] Loading weights as 4-bit NF4...") total_4bit = 0 total_bf16 = 0 replaced_linears = {} for shard_name in shards: shard_path = os.path.join(MODEL_PATH, shard_name) shard_keys = [k for k, v in weight_map.items() if v == shard_name] print(f" {shard_name} — {len(shard_keys)} tensors...", flush=True) with safe_open(shard_path, framework="pt", device="cpu") as f: for key in shard_keys: if key not in param_names and key not in buffer_names: continue tensor = f.get_tensor(key) if should_quantize(key) and key in param_names: owner, linear_attr = find_linear_owner(model, key) linear_id = id(getattr(owner, linear_attr)) if linear_id not in replaced_linears: old_linear = getattr(owner, linear_attr) in_f = old_linear.in_features out_f = old_linear.out_features has_bias = old_linear.bias is not None new_linear = bnb.nn.Linear4bit( in_f, out_f, bias=has_bias, quant_type="nf4", compute_dtype=torch.bfloat16, device=DEVICE, ) t = tensor.to(torch.bfloat16) new_linear.weight = bnb.nn.Params4bit( t, requires_grad=False, quant_type="nf4", ).cuda(0) if has_bias: new_linear.bias = None setattr(owner, linear_attr, new_linear) replaced_linears[id(new_linear)] = new_linear del old_linear, new_linear, t total_4bit += 1 else: parent, param_name = navigate_to_parent(model, key) t = tensor.to(torch.bfloat16).to(DEVICE) parent._parameters[param_name] = nn.Parameter(t, requires_grad=False) del t total_bf16 += 1 del tensor if (total_4bit + total_bf16) % 200 == 0: gc.collect() torch.cuda.empty_cache() gc.collect() torch.cuda.empty_cache() vram = torch.cuda.memory_allocated(0) / 1e9 print(f" {total_4bit}q+{total_bf16}b, VRAM: {vram:.1f} GB", flush=True) # Fix tied embeddings if model.lm_head.weight.device.type == 'meta': model.lm_head.weight = nn.Parameter( model.model.embed_tokens.weight.data.clone(), requires_grad=False ) # Init coding_gate and coding_norm 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) # Patch experts (SwiGLU clamp) patched = 0 for layer in model.model.layers: if not isinstance(layer, Fuse2AugmentedLayer): continue experts = getattr(layer, "experts", None) if experts is None: continue 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): gate = F.silu(g(x)) val = torch.clamp(gate * u(x), -lim, lim) return d(val) return forward expert.forward = make_fwd(gp, up, dp, SWIGLU_LIMIT) patched += 1 print(f"[patch] {patched} experts patched with SwiGLU clamp") # Router stability for layer in model.model.layers: if not isinstance(layer, Fuse2AugmentedLayer): continue router = layer.router top_k = router.top_k gate = router.gate def make_stable_fwd(g, tk): def forward(hidden_states): logits = g(hidden_states) scores = F.softplus(logits) scores = torch.clamp(scores, min=1e-6) scores = scores.sqrt() topk_weights, topk_indices = scores.topk(tk, dim=-1) topk_weights = topk_weights / (topk_weights.sum(dim=-1, keepdim=True) + 1e-8) return topk_weights, topk_indices, logits return forward router.forward = make_stable_fwd(gate, top_k) print(f"[patch] Router stability applied") model.set_coding_enabled(True) model.to(DEVICE) model.eval() vram = torch.cuda.memory_allocated(0) / 1e9 print(f"[load] Ready, VRAM: {vram:.1f} GB") return model, tok def generate(model, tok, prompt, max_new_tokens=1024, temperature=0.6, repetition_penalty=1.1): """Generate text from a prompt.""" # Format as chat messages = [{"role": "user", "content": prompt}] text = tok.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) print(f"[gen] Prompt: {prompt[:100]}...") print(f"[gen] Max tokens: {max_new_tokens}, temp: {temperature}") print(f"[gen] Generating...", flush=True) inputs = tok(text, return_tensors="pt").to(DEVICE) input_len = inputs["input_ids"].shape[1] t0 = time.time() with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=max_new_tokens, temperature=temperature, do_sample=temperature > 0, repetition_penalty=repetition_penalty, pad_token_id=tok.pad_token_id, use_cache=True, ) t1 = time.time() generated = tok.decode(outputs[0][input_len:], skip_special_tokens=True) elapsed = t1 - t0 tok_count = outputs.shape[1] - input_len tps = tok_count / elapsed print(f"[gen] {tok_count} tokens in {elapsed:.1f}s ({tps:.1f} tok/s)") return generated if __name__ == "__main__": prompt = sys.argv[1] if len(sys.argv) > 1 else "Write a Python fizzbuzz." model, tok = load_model_4bit() print("\n" + "=" * 60) print("GENERATING") print("=" * 60) result = generate(model, tok, prompt, max_new_tokens=2048, temperature=0.6) print("\n" + "=" * 60) print("OUTPUT") print("=" * 60) print(result)