| """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) |
| dev = "cuda" |
|
|
| |
| 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() |
|
|
| |
|
|
| 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 |
|
|
| |
| layer, method = build(ColumnParallelLinear, 3072, 9216, [3072, 3072, 3072], "col-parallel qkv") |
| |
| w_param = layer.weight |
| full = (torch.randn(9216, 3072, device=dev) * 0.02).to(torch.bfloat16) |
| w_param.weight_loader(w_param, full) |
| 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" |
|
|
| |
| 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() |
|
|
| |
| 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") |
|
|