|
| 1 | +# SPDX-License-Identifier: Apache-2.0 |
| 2 | +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project |
| 3 | +"""Tests for the per_token_group_fp8_quant helion kernel |
| 4 | +
|
| 5 | +Run `pytest tests/kernels/helion/test_per_token_group_fp8_quant.py`. |
| 6 | +""" |
| 7 | + |
| 8 | +from typing import Any |
| 9 | + |
| 10 | +import pytest |
| 11 | +import torch |
| 12 | +from torch._subclasses.fake_tensor import FakeTensorMode |
| 13 | + |
| 14 | +from tests.kernels.helion.utils import skip_if_platform_unsupported |
| 15 | +from tests.kernels.quant_utils import FP8_DTYPE |
| 16 | +from vllm.kernels.helion.case_key import CaseKey |
| 17 | +from vllm.kernels.helion.config_manager import ConfigManager |
| 18 | +from vllm.kernels.helion.ops.per_token_group_fp8_quant import ( |
| 19 | + _pick_cache, |
| 20 | + baseline, |
| 21 | + per_token_group_fp8_quant, |
| 22 | + pick_config, |
| 23 | +) |
| 24 | +from vllm.model_executor.layers.quantization.utils.quant_utils import ( |
| 25 | + get_fp8_min_max, |
| 26 | +) |
| 27 | +from vllm.utils.import_utils import has_helion |
| 28 | + |
| 29 | +if not has_helion(): |
| 30 | + pytest.skip( |
| 31 | + "Helion is not installed. Install with: pip install vllm[helion]", |
| 32 | + allow_module_level=True, |
| 33 | + ) |
| 34 | + |
| 35 | + |
| 36 | +def _generate_fake_input( |
| 37 | + num_tokens: int, hidden_size: int, group_size: int |
| 38 | +) -> tuple[Any, ...]: |
| 39 | + with FakeTensorMode(): |
| 40 | + input = torch.randn( |
| 41 | + (num_tokens, hidden_size), device="cuda", dtype=torch.bfloat16 |
| 42 | + ) |
| 43 | + output_q = torch.empty(input.shape, device=input.device, dtype=FP8_DTYPE) |
| 44 | + output_s = torch.empty( |
| 45 | + (num_tokens, hidden_size // group_size), |
| 46 | + device=input.device, |
| 47 | + dtype=torch.float32, |
| 48 | + ) |
| 49 | + use_ue8m0 = False |
| 50 | + column_major = False |
| 51 | + fp8_min, fp8_max = get_fp8_min_max() |
| 52 | + eps = 1e-10 |
| 53 | + args = ( |
| 54 | + input, |
| 55 | + output_q, |
| 56 | + output_s, |
| 57 | + group_size, |
| 58 | + eps, |
| 59 | + fp8_min, |
| 60 | + fp8_max, |
| 61 | + use_ue8m0, |
| 62 | + column_major, |
| 63 | + ) |
| 64 | + return args |
| 65 | + |
| 66 | + |
| 67 | +@pytest.fixture(autouse=True) |
| 68 | +def reset_config_manager_singleton(): |
| 69 | + ConfigManager.reset_instance() |
| 70 | + ConfigManager() |
| 71 | + yield |
| 72 | + ConfigManager.reset_instance() |
| 73 | + |
| 74 | + |
| 75 | +class TestPerTokenGroupFp8QuantConfigPicker: |
| 76 | + def setup_method(self): |
| 77 | + _pick_cache.clear() |
| 78 | + |
| 79 | + def test_config_picker_exact_match(self): |
| 80 | + config_keys = [ |
| 81 | + CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}), |
| 82 | + CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}), |
| 83 | + ] |
| 84 | + |
| 85 | + args = _generate_fake_input(16, 4096, 128) |
| 86 | + selected_key = pick_config(args, config_keys) |
| 87 | + assert selected_key == CaseKey( |
| 88 | + {"hidden_size": 4096, "group_size": 128, "num_tokens": 16} |
| 89 | + ) |
| 90 | + |
| 91 | + def test_config_picker_closest_match(self): |
| 92 | + config_keys = [ |
| 93 | + CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}), |
| 94 | + CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 32}), |
| 95 | + CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 16}), |
| 96 | + CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 32}), |
| 97 | + CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 16}), |
| 98 | + CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 32}), |
| 99 | + CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}), |
| 100 | + CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 32}), |
| 101 | + ] |
| 102 | + |
| 103 | + args = _generate_fake_input(20, 3000, 70) |
| 104 | + selected_key = pick_config(args, config_keys) |
| 105 | + assert selected_key == CaseKey( |
| 106 | + {"hidden_size": 2048, "group_size": 64, "num_tokens": 32} |
| 107 | + ) |
| 108 | + |
| 109 | + def test_config_picker_no_configs(self): |
| 110 | + config_keys: list[dict] = [] |
| 111 | + |
| 112 | + args = _generate_fake_input(16, 4096, 128) |
| 113 | + selected_key = pick_config(args, config_keys) |
| 114 | + assert selected_key is None |
| 115 | + |
| 116 | + def test_config_picker_fallback_to_largest(self): |
| 117 | + config_keys = [ |
| 118 | + CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}), |
| 119 | + CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 32}), |
| 120 | + CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 16}), |
| 121 | + CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 32}), |
| 122 | + CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 16}), |
| 123 | + CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 32}), |
| 124 | + CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}), |
| 125 | + CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 32}), |
| 126 | + ] |
| 127 | + |
| 128 | + args = _generate_fake_input(64, 8192, 256) |
| 129 | + selected_key = pick_config(args, config_keys) |
| 130 | + assert selected_key == CaseKey( |
| 131 | + {"hidden_size": 4096, "group_size": 128, "num_tokens": 32} |
| 132 | + ) |
| 133 | + |
| 134 | + |
| 135 | +class TestPerTokenGroupFp8QuantCorrectness: |
| 136 | + @pytest.mark.parametrize( |
| 137 | + "shape", [(31, 128), (32, 128), (63, 256), (64, 256), (16, 512), (2048, 5120)] |
| 138 | + ) |
| 139 | + @pytest.mark.parametrize("column_major", [False, True]) |
| 140 | + @pytest.mark.parametrize("tma_aligned", [False, True]) |
| 141 | + @pytest.mark.parametrize("scale_ue8m0", [False, True]) |
| 142 | + @pytest.mark.parametrize("group_size", [64, 128]) |
| 143 | + def test_per_token_group_fp8_quant( |
| 144 | + self, |
| 145 | + shape, |
| 146 | + column_major: bool, |
| 147 | + tma_aligned: bool, |
| 148 | + scale_ue8m0: bool, |
| 149 | + group_size: int, |
| 150 | + ): |
| 151 | + skip_if_platform_unsupported("per_token_group_fp8_quant") |
| 152 | + |
| 153 | + torch.manual_seed(42) |
| 154 | + num_tokens, hidden_size = shape |
| 155 | + fp8_min, fp8_max = get_fp8_min_max() |
| 156 | + eps = 1e-10 |
| 157 | + input = ( |
| 158 | + torch.randn((num_tokens, hidden_size), device="cuda", dtype=torch.bfloat16) |
| 159 | + * 8 |
| 160 | + ) |
| 161 | + ref_q = torch.empty(input.shape, device=input.device, dtype=FP8_DTYPE) |
| 162 | + ops_q = ref_q.clone() |
| 163 | + |
| 164 | + groups_per_row = hidden_size // group_size |
| 165 | + if column_major: |
| 166 | + if tma_aligned: |
| 167 | + tma_alignment = 4 |
| 168 | + tma_aligned_m = ( |
| 169 | + (num_tokens + tma_alignment - 1) // tma_alignment * tma_alignment |
| 170 | + ) |
| 171 | + shape = (num_tokens, groups_per_row) |
| 172 | + stride = (1, tma_aligned_m) |
| 173 | + ref_s = torch.empty_strided( |
| 174 | + shape, stride, device=input.device, dtype=torch.float32 |
| 175 | + ) |
| 176 | + else: |
| 177 | + ref_s = torch.empty( |
| 178 | + (groups_per_row, num_tokens), |
| 179 | + device=input.device, |
| 180 | + dtype=torch.float32, |
| 181 | + ).transpose(0, 1) |
| 182 | + else: |
| 183 | + ref_s = torch.empty( |
| 184 | + (num_tokens, groups_per_row), device=input.device, dtype=torch.float32 |
| 185 | + ) |
| 186 | + |
| 187 | + ops_s = ref_s.clone() |
| 188 | + |
| 189 | + baseline( |
| 190 | + input, |
| 191 | + ref_q, |
| 192 | + ref_s, |
| 193 | + group_size, |
| 194 | + eps, |
| 195 | + fp8_min, |
| 196 | + fp8_max, |
| 197 | + scale_ue8m0, |
| 198 | + column_major, |
| 199 | + tma_aligned, |
| 200 | + ) |
| 201 | + per_token_group_fp8_quant( |
| 202 | + input, |
| 203 | + ops_q, |
| 204 | + ops_s, |
| 205 | + group_size, |
| 206 | + eps, |
| 207 | + fp8_min, |
| 208 | + fp8_max, |
| 209 | + scale_ue8m0, |
| 210 | + column_major, |
| 211 | + tma_aligned, |
| 212 | + ) |
| 213 | + |
| 214 | + assert torch.allclose(ref_s, ops_s) |
| 215 | + # allow 1 ULP difference |
| 216 | + assert ( |
| 217 | + ref_q.view(torch.uint8).to(torch.int16) |
| 218 | + - ops_q.view(torch.uint8).to(torch.int16) |
| 219 | + ).abs().max() <= 1 |
| 220 | + |
| 221 | + |
| 222 | +class TestPerTokenGroupFp8QuantIntegration: |
| 223 | + def test_kernel_registration_integration(self): |
| 224 | + from vllm.kernels.helion.register import get_registered_kernels |
| 225 | + |
| 226 | + registered_kernels = get_registered_kernels() |
| 227 | + assert "per_token_group_fp8_quant" in registered_kernels |
| 228 | + |
| 229 | + kernel_wrapper = registered_kernels["per_token_group_fp8_quant"] |
| 230 | + assert kernel_wrapper.op_name == "per_token_group_fp8_quant" |
| 231 | + assert kernel_wrapper._config_picker is not None |
| 232 | + assert kernel_wrapper._mutates_args == ["output_q", "output_s"] |
| 233 | + |
| 234 | + def test_fake_impl_functionality(self): |
| 235 | + skip_if_platform_unsupported("per_token_group_fp8_quant") |
| 236 | + from vllm.kernels.helion.register import get_registered_kernels |
| 237 | + |
| 238 | + registered_kernels = get_registered_kernels() |
| 239 | + kernel_wrapper = registered_kernels["per_token_group_fp8_quant"] |
| 240 | + fake_impl = kernel_wrapper._fake_impl |
| 241 | + |
| 242 | + args = _generate_fake_input(16, 4096, 128) |
| 243 | + assert fake_impl(*args) is None |
0 commit comments