mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-17 11:10:16 +00:00
[Kernel][Helion][1/N] Add Helion kernel for per_token_group_fp8_quant (#36902)
Signed-off-by: Sean Chen <[email protected]> Co-authored-by: Yanan Cao <[email protected]> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
Yanan Cao
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
79f8c5bd8c
commit
2ec6594db9
@@ -398,7 +398,7 @@ steps:
|
||||
- tests/kernels/helion/
|
||||
- vllm/platforms/rocm.py
|
||||
commands:
|
||||
- pip install helion==1.0.0
|
||||
- pip install helion==1.1.0
|
||||
- pytest -v -s kernels/helion/
|
||||
|
||||
- label: Kernels Mamba Test # TBD
|
||||
|
||||
@@ -237,7 +237,7 @@ steps:
|
||||
- vllm/utils/import_utils.py
|
||||
- tests/kernels/helion/
|
||||
commands:
|
||||
- pip install helion==1.0.0
|
||||
- pip install helion==1.1.0
|
||||
- pytest -v -s kernels/helion/
|
||||
|
||||
|
||||
|
||||
@@ -1229,7 +1229,7 @@ setup(
|
||||
# NOTE: When updating helion version, also update CI files:
|
||||
# - .buildkite/test_areas/kernels.yaml
|
||||
# - .buildkite/test-amd.yaml
|
||||
"helion": ["helion==1.0.0"],
|
||||
"helion": ["helion==1.1.0"],
|
||||
# Optional deps for gRPC server (vllm serve --grpc)
|
||||
"grpc": ["smg-grpc-servicer[vllm] >= 0.5.2"],
|
||||
# Optional deps for OpenTelemetry tracing
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests for the per_token_group_fp8_quant helion kernel
|
||||
|
||||
Run `pytest tests/kernels/helion/test_per_token_group_fp8_quant.py`.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch._subclasses.fake_tensor import FakeTensorMode
|
||||
|
||||
from tests.kernels.helion.utils import skip_if_platform_unsupported
|
||||
from tests.kernels.quant_utils import FP8_DTYPE
|
||||
from vllm.kernels.helion.case_key import CaseKey
|
||||
from vllm.kernels.helion.config_manager import ConfigManager
|
||||
from vllm.kernels.helion.ops.per_token_group_fp8_quant import (
|
||||
_pick_cache,
|
||||
baseline,
|
||||
per_token_group_fp8_quant,
|
||||
pick_config,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
get_fp8_min_max,
|
||||
)
|
||||
from vllm.utils.import_utils import has_helion
|
||||
|
||||
if not has_helion():
|
||||
pytest.skip(
|
||||
"Helion is not installed. Install with: pip install vllm[helion]",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
|
||||
def _generate_fake_input(
|
||||
num_tokens: int, hidden_size: int, group_size: int
|
||||
) -> tuple[Any, ...]:
|
||||
with FakeTensorMode():
|
||||
input = torch.randn(
|
||||
(num_tokens, hidden_size), device="cuda", dtype=torch.bfloat16
|
||||
)
|
||||
output_q = torch.empty(input.shape, device=input.device, dtype=FP8_DTYPE)
|
||||
output_s = torch.empty(
|
||||
(num_tokens, hidden_size // group_size),
|
||||
device=input.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
use_ue8m0 = False
|
||||
column_major = False
|
||||
fp8_min, fp8_max = get_fp8_min_max()
|
||||
eps = 1e-10
|
||||
args = (
|
||||
input,
|
||||
output_q,
|
||||
output_s,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
use_ue8m0,
|
||||
column_major,
|
||||
)
|
||||
return args
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_config_manager_singleton():
|
||||
ConfigManager.reset_instance()
|
||||
ConfigManager()
|
||||
yield
|
||||
ConfigManager.reset_instance()
|
||||
|
||||
|
||||
class TestPerTokenGroupFp8QuantConfigPicker:
|
||||
def setup_method(self):
|
||||
_pick_cache.clear()
|
||||
|
||||
def test_config_picker_exact_match(self):
|
||||
config_keys = [
|
||||
CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}),
|
||||
]
|
||||
|
||||
args = _generate_fake_input(16, 4096, 128)
|
||||
selected_key = pick_config(args, config_keys)
|
||||
assert selected_key == CaseKey(
|
||||
{"hidden_size": 4096, "group_size": 128, "num_tokens": 16}
|
||||
)
|
||||
|
||||
def test_config_picker_closest_match(self):
|
||||
config_keys = [
|
||||
CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 32}),
|
||||
]
|
||||
|
||||
args = _generate_fake_input(20, 3000, 70)
|
||||
selected_key = pick_config(args, config_keys)
|
||||
assert selected_key == CaseKey(
|
||||
{"hidden_size": 2048, "group_size": 64, "num_tokens": 32}
|
||||
)
|
||||
|
||||
def test_config_picker_no_configs(self):
|
||||
config_keys: list[dict] = []
|
||||
|
||||
args = _generate_fake_input(16, 4096, 128)
|
||||
selected_key = pick_config(args, config_keys)
|
||||
assert selected_key is None
|
||||
|
||||
def test_config_picker_fallback_to_largest(self):
|
||||
config_keys = [
|
||||
CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 64, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 2048, "group_size": 128, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 64, "num_tokens": 32}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 16}),
|
||||
CaseKey({"hidden_size": 4096, "group_size": 128, "num_tokens": 32}),
|
||||
]
|
||||
|
||||
args = _generate_fake_input(64, 8192, 256)
|
||||
selected_key = pick_config(args, config_keys)
|
||||
assert selected_key == CaseKey(
|
||||
{"hidden_size": 4096, "group_size": 128, "num_tokens": 32}
|
||||
)
|
||||
|
||||
|
||||
class TestPerTokenGroupFp8QuantCorrectness:
|
||||
@pytest.mark.parametrize(
|
||||
"shape", [(31, 128), (32, 128), (63, 256), (64, 256), (16, 512), (2048, 5120)]
|
||||
)
|
||||
@pytest.mark.parametrize("column_major", [False, True])
|
||||
@pytest.mark.parametrize("tma_aligned", [False, True])
|
||||
@pytest.mark.parametrize("scale_ue8m0", [False, True])
|
||||
@pytest.mark.parametrize("group_size", [64, 128])
|
||||
def test_per_token_group_fp8_quant(
|
||||
self,
|
||||
shape,
|
||||
column_major: bool,
|
||||
tma_aligned: bool,
|
||||
scale_ue8m0: bool,
|
||||
group_size: int,
|
||||
):
|
||||
skip_if_platform_unsupported("per_token_group_fp8_quant")
|
||||
|
||||
torch.manual_seed(42)
|
||||
num_tokens, hidden_size = shape
|
||||
fp8_min, fp8_max = get_fp8_min_max()
|
||||
eps = 1e-10
|
||||
input = (
|
||||
torch.randn((num_tokens, hidden_size), device="cuda", dtype=torch.bfloat16)
|
||||
* 8
|
||||
)
|
||||
ref_q = torch.empty(input.shape, device=input.device, dtype=FP8_DTYPE)
|
||||
ops_q = ref_q.clone()
|
||||
|
||||
groups_per_row = hidden_size // group_size
|
||||
if column_major:
|
||||
if tma_aligned:
|
||||
tma_alignment = 4
|
||||
tma_aligned_m = (
|
||||
(num_tokens + tma_alignment - 1) // tma_alignment * tma_alignment
|
||||
)
|
||||
shape = (num_tokens, groups_per_row)
|
||||
stride = (1, tma_aligned_m)
|
||||
ref_s = torch.empty_strided(
|
||||
shape, stride, device=input.device, dtype=torch.float32
|
||||
)
|
||||
else:
|
||||
ref_s = torch.empty(
|
||||
(groups_per_row, num_tokens),
|
||||
device=input.device,
|
||||
dtype=torch.float32,
|
||||
).transpose(0, 1)
|
||||
else:
|
||||
ref_s = torch.empty(
|
||||
(num_tokens, groups_per_row), device=input.device, dtype=torch.float32
|
||||
)
|
||||
|
||||
ops_s = ref_s.clone()
|
||||
|
||||
baseline(
|
||||
input,
|
||||
ref_q,
|
||||
ref_s,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
scale_ue8m0,
|
||||
column_major,
|
||||
tma_aligned,
|
||||
)
|
||||
per_token_group_fp8_quant(
|
||||
input,
|
||||
ops_q,
|
||||
ops_s,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
scale_ue8m0,
|
||||
column_major,
|
||||
tma_aligned,
|
||||
)
|
||||
|
||||
assert torch.allclose(ref_s, ops_s)
|
||||
# allow 1 ULP difference
|
||||
assert (
|
||||
ref_q.view(torch.uint8).to(torch.int16)
|
||||
- ops_q.view(torch.uint8).to(torch.int16)
|
||||
).abs().max() <= 1
|
||||
|
||||
|
||||
class TestPerTokenGroupFp8QuantIntegration:
|
||||
def test_kernel_registration_integration(self):
|
||||
from vllm.kernels.helion.register import get_registered_kernels
|
||||
|
||||
registered_kernels = get_registered_kernels()
|
||||
assert "per_token_group_fp8_quant" in registered_kernels
|
||||
|
||||
kernel_wrapper = registered_kernels["per_token_group_fp8_quant"]
|
||||
assert kernel_wrapper.op_name == "per_token_group_fp8_quant"
|
||||
assert kernel_wrapper._config_picker is not None
|
||||
assert kernel_wrapper._mutates_args == ["output_q", "output_s"]
|
||||
|
||||
def test_fake_impl_functionality(self):
|
||||
skip_if_platform_unsupported("per_token_group_fp8_quant")
|
||||
from vllm.kernels.helion.register import get_registered_kernels
|
||||
|
||||
registered_kernels = get_registered_kernels()
|
||||
kernel_wrapper = registered_kernels["per_token_group_fp8_quant"]
|
||||
fake_impl = kernel_wrapper._fake_impl
|
||||
|
||||
args = _generate_fake_input(16, 4096, 128)
|
||||
assert fake_impl(*args) is None
|
||||
@@ -713,6 +713,7 @@ class TestHelionKernelWrapper:
|
||||
|
||||
new_op = Mock()
|
||||
registered_ops: dict[str, Mock] = {}
|
||||
mutates_args = ["y"]
|
||||
|
||||
class MockNamespace:
|
||||
def __getattr__(self, name):
|
||||
@@ -748,6 +749,7 @@ class TestHelionKernelWrapper:
|
||||
raw_kernel_func=sample_kernel,
|
||||
op_name="test_kernel",
|
||||
fake_impl=fake_impl,
|
||||
mutates_args=mutates_args,
|
||||
config_picker=default_picker,
|
||||
)
|
||||
result = wrapper._get_or_register_custom_op()
|
||||
@@ -755,6 +757,7 @@ class TestHelionKernelWrapper:
|
||||
mock_register.assert_called_once()
|
||||
assert result is new_op
|
||||
assert mock_register.call_args[1]["op_func"] is mock_decorated
|
||||
assert mock_register.call_args[1]["mutates_args"] is mutates_args
|
||||
|
||||
|
||||
class TestKernelRegistry:
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Helion Kernel test utils"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm.kernels.helion.config_manager import ConfigManager
|
||||
|
||||
|
||||
def skip_if_platform_unsupported(op_name: str):
|
||||
try:
|
||||
from vllm.kernels.helion.utils import get_canonical_gpu_name
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA not available")
|
||||
|
||||
platform = get_canonical_gpu_name()
|
||||
|
||||
try:
|
||||
config_manager = ConfigManager.get_instance()
|
||||
except RuntimeError:
|
||||
config_manager = ConfigManager()
|
||||
|
||||
configs = config_manager.get_platform_configs(op_name, platform)
|
||||
if len(configs) == 0:
|
||||
pytest.skip(f"Current GPU platform not supported for {op_name} kernel")
|
||||
|
||||
except (ImportError, RuntimeError, KeyError):
|
||||
pytest.skip(f"Error detecting platform support for {op_name} kernel")
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,232 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from itertools import product
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.kernels.helion.case_key import CaseKey
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
get_fp8_min_max,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.import_utils import has_helion
|
||||
|
||||
if not has_helion():
|
||||
raise ImportError(
|
||||
"Helion kernel requires helion to be installed. "
|
||||
"Install it with: pip install helion"
|
||||
)
|
||||
|
||||
import helion
|
||||
import helion.language as hl
|
||||
|
||||
from vllm.kernels.helion.register import register_kernel
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def generate_inputs() -> dict[CaseKey, tuple[Any, ...]]:
|
||||
# TODO(xiaohongchen1991): it is difficult for kernel author to cover all
|
||||
# input property combination. Currently, dtypes are fixed. We need
|
||||
# optimization to bucket/skip some combinations
|
||||
num_tokens_list = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192]
|
||||
hidden_size_list = [2048, 4096, 5120]
|
||||
group_size_list = [128]
|
||||
in_dtype: torch.dtype = torch.bfloat16
|
||||
out_dtype: torch.dtype = current_platform.fp8_dtype()
|
||||
scale_dtype: torch.dtype = torch.float32
|
||||
|
||||
use_ue8m0 = False
|
||||
column_major = False
|
||||
fp8_min, fp8_max = get_fp8_min_max()
|
||||
eps = 1e-10
|
||||
|
||||
inputs = {}
|
||||
|
||||
for hidden_size, group_size, num_tokens in product(
|
||||
hidden_size_list, group_size_list, num_tokens_list
|
||||
):
|
||||
input = torch.randn(num_tokens, hidden_size, device="cuda", dtype=in_dtype)
|
||||
output_q = torch.empty(input.shape, device=input.device, dtype=out_dtype)
|
||||
output_s = torch.empty(
|
||||
(num_tokens, hidden_size // group_size),
|
||||
device=input.device,
|
||||
dtype=scale_dtype,
|
||||
)
|
||||
config_key = CaseKey(
|
||||
{
|
||||
"hidden_size": hidden_size,
|
||||
"group_size": group_size,
|
||||
"num_tokens": num_tokens,
|
||||
}
|
||||
)
|
||||
inputs[config_key] = (
|
||||
input,
|
||||
output_q,
|
||||
output_s,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
use_ue8m0,
|
||||
column_major,
|
||||
False,
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
|
||||
_pick_cache: dict[tuple[int, int, int], CaseKey | None] = {}
|
||||
|
||||
|
||||
def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | None:
|
||||
"""Pick the best pre-tuned config for the given input shape.
|
||||
|
||||
Selection strategy:
|
||||
1. Find the closest hidden_size among available configs
|
||||
(exact match preferred).
|
||||
2. Find the closest group_size among available configs
|
||||
(exact match preferred).
|
||||
3. Among the num_tokens values tuned for that hidden_size and group_size, pick
|
||||
the smallest num_tokens >= the input's num_tokens. If the input is
|
||||
larger than all available num_tokens, fall back to the largest.
|
||||
"""
|
||||
|
||||
if not config_keys:
|
||||
return None
|
||||
|
||||
input, _, _, group_size, *_ = args
|
||||
num_tokens, hidden_size = input.shape
|
||||
|
||||
cache_key = (num_tokens, group_size, hidden_size)
|
||||
cached = _pick_cache.get(cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
configs: dict[int, dict[int, list[int]]] = {}
|
||||
for key in config_keys:
|
||||
if key.is_default():
|
||||
continue
|
||||
configs.setdefault(key["hidden_size"], {}).setdefault(
|
||||
key["group_size"], []
|
||||
).append(key["num_tokens"])
|
||||
|
||||
if not configs:
|
||||
return None
|
||||
|
||||
best_hidden_size = min(configs, key=lambda s: abs(s - hidden_size))
|
||||
best_group_size = min(configs[best_hidden_size], key=lambda s: abs(s - group_size))
|
||||
available_num_tokens = sorted(configs[best_hidden_size][best_group_size])
|
||||
best_num_tokens = next(
|
||||
(n for n in available_num_tokens if n >= num_tokens), available_num_tokens[-1]
|
||||
)
|
||||
|
||||
result = CaseKey(
|
||||
{
|
||||
"hidden_size": best_hidden_size,
|
||||
"group_size": best_group_size,
|
||||
"num_tokens": best_num_tokens,
|
||||
}
|
||||
)
|
||||
_pick_cache[cache_key] = result
|
||||
return result
|
||||
|
||||
|
||||
def fake_impl(
|
||||
input: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_q: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_s: torch.Tensor, # [num_tokens, groups_per_row]
|
||||
group_size: int,
|
||||
eps: float,
|
||||
fp8_min: float,
|
||||
fp8_max: float,
|
||||
scale_ue8m0: bool,
|
||||
# Unused dummy args
|
||||
# Kept for consistency with existing kernel interface
|
||||
dummy_is_scale_transposed: bool = False,
|
||||
dummy_is_tma_aligned: bool = False,
|
||||
) -> None:
|
||||
return
|
||||
|
||||
|
||||
def baseline(
|
||||
input: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_q: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_s: torch.Tensor, # [num_tokens, groups_per_row]
|
||||
group_size: int,
|
||||
eps: float,
|
||||
fp8_min: float,
|
||||
fp8_max: float,
|
||||
scale_ue8m0: bool,
|
||||
dummy_is_scale_transposed: bool = False,
|
||||
dummy_is_tma_aligned: bool = False,
|
||||
) -> None:
|
||||
torch.ops._C.per_token_group_fp8_quant(
|
||||
input,
|
||||
output_q,
|
||||
output_s,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
scale_ue8m0,
|
||||
dummy_is_scale_transposed,
|
||||
dummy_is_tma_aligned,
|
||||
)
|
||||
|
||||
|
||||
@register_kernel(
|
||||
mutates_args=["output_q", "output_s"],
|
||||
config_picker=pick_config,
|
||||
input_generator=generate_inputs,
|
||||
fake_impl=fake_impl,
|
||||
helion_settings=helion.Settings(
|
||||
autotune_baseline_fn=baseline,
|
||||
),
|
||||
) # type: ignore[misc]
|
||||
def per_token_group_fp8_quant(
|
||||
input: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_q: torch.Tensor, # [num_tokens, hidden_size]
|
||||
output_s: torch.Tensor, # [num_tokens, groups_per_row]
|
||||
group_size: int,
|
||||
eps: float,
|
||||
fp8_min: float,
|
||||
fp8_max: float,
|
||||
scale_ue8m0: bool,
|
||||
# Unused dummy args
|
||||
# Kept for consistency with existing kernel interface
|
||||
dummy_is_scale_transposed: bool = False,
|
||||
dummy_is_tma_aligned: bool = False,
|
||||
) -> None:
|
||||
# This code assumes batch_dim and num_tokens are flattened
|
||||
assert input.ndim == 2
|
||||
num_tokens, hidden_size = input.shape
|
||||
hl.specialize(hidden_size)
|
||||
hl.specialize(group_size)
|
||||
|
||||
groups_per_row = output_s.shape[1]
|
||||
hl.specialize(groups_per_row)
|
||||
assert hidden_size % group_size == 0 and hidden_size // group_size == groups_per_row
|
||||
assert output_s.ndim == 2 and output_s.dtype == torch.float32
|
||||
|
||||
input = input.view(num_tokens, -1, group_size)
|
||||
output_q = output_q.view(num_tokens, -1, group_size)
|
||||
for tile_m, tile_gn, tile_n in hl.tile(
|
||||
[num_tokens, groups_per_row, group_size], block_size=[1, None, group_size]
|
||||
):
|
||||
x_blk = input[tile_m, tile_gn, tile_n]
|
||||
y_s_blk = torch.clamp(torch.amax(torch.abs(x_blk), dim=-1), min=eps)
|
||||
y_s_blk = y_s_blk / fp8_max
|
||||
|
||||
if scale_ue8m0:
|
||||
y_s_blk = torch.exp2(torch.ceil(torch.log2(y_s_blk)))
|
||||
|
||||
y_q_blk = torch.clamp(x_blk / y_s_blk[:, :, None], fp8_min, fp8_max).to(
|
||||
output_q.dtype
|
||||
)
|
||||
|
||||
output_s[tile_m, tile_gn] = y_s_blk
|
||||
output_q[tile_m, tile_gn, tile_n] = y_q_blk
|
||||
@@ -260,6 +260,7 @@ class HelionKernelWrapper:
|
||||
op_name: str,
|
||||
fake_impl: Callable,
|
||||
config_picker: ConfigPicker,
|
||||
mutates_args: list[str] | None = None,
|
||||
helion_settings: helion.Settings | None = None,
|
||||
input_generator: (Callable[[], dict[CaseKey, tuple[Any, ...]]] | None) = None,
|
||||
):
|
||||
@@ -272,6 +273,7 @@ class HelionKernelWrapper:
|
||||
self.helion_settings = helion_settings
|
||||
self._config_picker = config_picker
|
||||
self._input_generator = input_generator
|
||||
self._mutates_args = mutates_args
|
||||
self._configured_kernel: ConfiguredHelionKernel | None = None
|
||||
# TODO(@gmagogsfm): Remove this disable flag once integrated with vLLM IR,
|
||||
# which handles op enablement/disablement.
|
||||
@@ -357,7 +359,7 @@ class HelionKernelWrapper:
|
||||
direct_register_custom_op(
|
||||
op_name=self.op_name,
|
||||
op_func=configured_kernel._decorated_kernel,
|
||||
mutates_args=None,
|
||||
mutates_args=self._mutates_args,
|
||||
fake_impl=self._fake_impl,
|
||||
target_lib=vllm_helion_lib,
|
||||
)
|
||||
@@ -402,6 +404,7 @@ def register_kernel(
|
||||
*,
|
||||
config_picker: ConfigPicker,
|
||||
fake_impl: Callable | None = None,
|
||||
mutates_args: list[str] | None = None,
|
||||
helion_settings: helion.Settings | None = None,
|
||||
input_generator: (Callable[[], dict[CaseKey, tuple[Any, ...]]] | None) = None,
|
||||
) -> Callable[[Callable], HelionKernelWrapper]:
|
||||
@@ -455,6 +458,7 @@ def register_kernel(
|
||||
op_name=final_op_name,
|
||||
fake_impl=final_fake_impl,
|
||||
config_picker=config_picker,
|
||||
mutates_args=mutates_args,
|
||||
helion_settings=helion_settings,
|
||||
input_generator=input_generator,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user