Skip to content

Commit 2ec6594

Browse files
xiaohongchen1991gmagogsfmmergify[bot]
authored
[Kernel][Helion][1/N] Add Helion kernel for per_token_group_fp8_quant (vllm-project#36902)
Signed-off-by: Sean Chen <seachen@redhat.com> Co-authored-by: Yanan Cao <gmagogsfm@gmail.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
1 parent 79f8c5b commit 2ec6594

10 files changed

Lines changed: 4347 additions & 4 deletions

File tree

.buildkite/test-amd.yaml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -398,7 +398,7 @@ steps:
398398
- tests/kernels/helion/
399399
- vllm/platforms/rocm.py
400400
commands:
401-
- pip install helion==1.0.0
401+
- pip install helion==1.1.0
402402
- pytest -v -s kernels/helion/
403403

404404
- label: Kernels Mamba Test # TBD

.buildkite/test_areas/kernels.yaml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -237,7 +237,7 @@ steps:
237237
- vllm/utils/import_utils.py
238238
- tests/kernels/helion/
239239
commands:
240-
- pip install helion==1.0.0
240+
- pip install helion==1.1.0
241241
- pytest -v -s kernels/helion/
242242

243243

setup.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1229,7 +1229,7 @@ def add_vllm_package_data(filename: str) -> None:
12291229
# NOTE: When updating helion version, also update CI files:
12301230
# - .buildkite/test_areas/kernels.yaml
12311231
# - .buildkite/test-amd.yaml
1232-
"helion": ["helion==1.0.0"],
1232+
"helion": ["helion==1.1.0"],
12331233
# Optional deps for gRPC server (vllm serve --grpc)
12341234
"grpc": ["smg-grpc-servicer[vllm] >= 0.5.2"],
12351235
# Optional deps for OpenTelemetry tracing
Lines changed: 243 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,243 @@
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

tests/kernels/helion/test_register.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -713,6 +713,7 @@ def default_picker(args, config_keys):
713713

714714
new_op = Mock()
715715
registered_ops: dict[str, Mock] = {}
716+
mutates_args = ["y"]
716717

717718
class MockNamespace:
718719
def __getattr__(self, name):
@@ -748,13 +749,15 @@ def register_side_effect(op_name, op_func, **kwargs):
748749
raw_kernel_func=sample_kernel,
749750
op_name="test_kernel",
750751
fake_impl=fake_impl,
752+
mutates_args=mutates_args,
751753
config_picker=default_picker,
752754
)
753755
result = wrapper._get_or_register_custom_op()
754756

755757
mock_register.assert_called_once()
756758
assert result is new_op
757759
assert mock_register.call_args[1]["op_func"] is mock_decorated
760+
assert mock_register.call_args[1]["mutates_args"] is mutates_args
758761

759762

760763
class TestKernelRegistry:

tests/kernels/helion/utils.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
# SPDX-License-Identifier: Apache-2.0
2+
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3+
"""Helion Kernel test utils"""
4+
5+
import pytest
6+
import torch
7+
8+
from vllm.kernels.helion.config_manager import ConfigManager
9+
10+
11+
def skip_if_platform_unsupported(op_name: str):
12+
try:
13+
from vllm.kernels.helion.utils import get_canonical_gpu_name
14+
15+
if not torch.cuda.is_available():
16+
pytest.skip("CUDA not available")
17+
18+
platform = get_canonical_gpu_name()
19+
20+
try:
21+
config_manager = ConfigManager.get_instance()
22+
except RuntimeError:
23+
config_manager = ConfigManager()
24+
25+
configs = config_manager.get_platform_configs(op_name, platform)
26+
if len(configs) == 0:
27+
pytest.skip(f"Current GPU platform not supported for {op_name} kernel")
28+
29+
except (ImportError, RuntimeError, KeyError):
30+
pytest.skip(f"Error detecting platform support for {op_name} kernel")

0 commit comments

Comments
 (0)