"""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")