svd-wan / test_svd_layer.py
Yi30's picture
Upload folder using huggingface_hub
312d5ce verified
Raw
History Blame Contribute Delete
3.57 kB
"""Verify quark_svdquant online W4A8 SVD on a real vllm parallel linear (CUDA)."""
import os
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "2")
import torch
from vllm.model_executor.layers.linear import ColumnParallelLinear, RowParallelLinear
from vllm_omni.quantization.quark_w4a8_config import DiffusionQuarkW4A8Config, QuarkW4A8SVDLinearMethod, QuarkW4A8LinearMethod
torch.manual_seed(0)
torch.set_default_dtype(torch.bfloat16) # as the diffusion model runner does
dev = "cuda"
# Minimal TP group stub for standalone layer testing (tp=1 semantics).
from vllm.distributed import parallel_state as _ps
class _FakeGroup:
rank_in_group = 0
world_size = 1
def rank(self):
return 0
_ps._TP = _FakeGroup()
_ps._PP = _FakeGroup()
_ps._DP = _FakeGroup()
# single-process TP=1 init so the parallel linear layers can be built standalone
def build(layer_cls, in_size, out_size, out_parts, cls_name):
qc = DiffusionQuarkW4A8Config(svd_rank=32)
layer = layer_cls(in_size, out_size, bias=False, quant_config=qc, disable_tp=True)
method = layer.quant_method
print(f"{cls_name}: method={type(method).__name__} derive_factors={getattr(method,'derive_factors',None)}")
assert isinstance(method, QuarkW4A8SVDLinearMethod), f"expected SVD method, got {type(method)}"
return layer, method
# --- column parallel (to_qkv-like): out 9216 = 3 x 3072, in 3072
layer, method = build(ColumnParallelLinear, 3072, 9216, [3072, 3072, 3072], "col-parallel qkv")
# load a random bf16 weight via the param's weight_loader (shard 0 only, tp=1 -> full)
w_param = layer.weight
full = (torch.randn(9216, 3072, device=dev) * 0.02).to(torch.bfloat16)
w_param.weight_loader(w_param, full) # tp=1: full fused matrix
ref = full.float()
layer = layer.to(dev)
method.process_weights_after_loading(layer)
print(" buffers after process:")
for name, buf in list(layer.named_buffers()) + [(n, p) for n, p in layer.named_parameters() if p is not None]:
t = buf if isinstance(buf, torch.Tensor) else None
if t is None:
continue
kind = "buffer" if name in layer._buffers else "param"
print(f" {kind:6s} {name:16s} {str(t.dtype):16s} {tuple(t.shape)}")
assert isinstance(getattr(layer, "_kernel_weight", None), torch.Tensor)
assert layer._kernel_weight.dtype == torch.uint8
assert layer._kernel_weight.shape == (9216, 3072 // 2), layer._kernel_weight.shape
assert layer._kernel_scale.shape == (9216, 3072 // 32)
assert layer.proj_up.shape == (9216, 32) and layer.proj_down.shape == (32, 3072)
assert not hasattr(layer, "weight") or layer.weight is None, "bf16 weight should be dropped"
# forward vs bf16 reference
x = torch.randn(64, 3072, device=dev, dtype=torch.bfloat16)
out, _ = layer(x)
ref_out = x.float() @ ref.t()
err_svd = (out.float() - ref_out).norm() / ref_out.norm()
# --- plain W4A8 (no svd) on a fresh layer for comparison
qc_plain = DiffusionQuarkW4A8Config(svd_rank=None)
layer2 = ColumnParallelLinear(3072, 9216, bias=False, quant_config=qc_plain, disable_tp=True)
m2 = layer2.quant_method
assert isinstance(m2, QuarkW4A8LinearMethod)
w2 = layer2.weight
w2.weight_loader(w2, full)
layer2 = layer2.to(dev)
m2.process_weights_after_loading(layer2)
out2, _ = layer2(x)
err_plain = (out2.float() - ref_out).norm() / ref_out.norm()
print(f" plain W4A8 rel err vs bf16: {err_plain.item():.4f}")
print(f" SVD W4A8 rel err vs bf16: {err_svd.item():.4f}")
assert err_svd < err_plain, "SVD correction should reduce error vs plain"
print("LAYER TEST PASSED: 4-bit buffers resident, SVD factors derived, forward numerics verified")