Mini-Whale-1-12B / fuse2_qlora_infer.py
Akahsizrr's picture
Upload fuse2_qlora_infer.py
2dc27e3 verified
Raw History Blame Contribute Delete
9.39 kB
"""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)