[AMD] Enable ROCm kvcache JIT path and add AMD CI coverage. (#18992)

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Hubert Lu
2026-02-24 22:15:05 -08:00
committed by GitHub
parent aff2f130ec
commit 8bd644765f
10 changed files with 168 additions and 12 deletions

View File

@@ -61,6 +61,9 @@ KERNEL_PATH = _resolve_kernel_path()
DEFAULT_INCLUDE = [str(KERNEL_PATH / "include")]
DEFAULT_CFLAGS = ["-std=c++20", "-O3"]
DEFAULT_CUDA_CFLAGS = ["-std=c++20", "-O3", "--expt-relaxed-constexpr"]
DEFAULT_HIP_CFLAGS = [
flag for flag in DEFAULT_CUDA_CFLAGS if flag != "--expt-relaxed-constexpr"
]
DEFAULT_LDFLAGS = []
CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, bool, torch.dtype]
@@ -77,6 +80,12 @@ CPP_DTYPE_MAP = {
}
# AMD/ROCm note:
@cache_once
def is_hip_runtime() -> bool:
return bool(torch.version.hip)
def make_cpp_args(*args: CPP_TEMPLATE_TYPE) -> CPPArgList:
def _convert(arg: CPP_TEMPLATE_TYPE) -> str:
if isinstance(arg, bool):
@@ -156,6 +165,10 @@ def load_jit(
# Override TVM_FFI_CUDA_ARCH_LIST if it does not exist.
env_key = "TVM_FFI_CUDA_ARCH_LIST"
env_existed = env_key in os.environ
selected_cuda_cflags = DEFAULT_CUDA_CFLAGS
if is_hip_runtime():
selected_cuda_cflags = DEFAULT_HIP_CFLAGS
extra_cuda_cflags = ["-DUSE_ROCM"] + extra_cuda_cflags
if not env_existed:
os.environ[env_key] = _get_cuda_arch_list()
try:
@@ -164,7 +177,7 @@ def load_jit(
cpp_sources=cpp_sources,
cuda_sources=cuda_sources,
extra_cflags=DEFAULT_CFLAGS + extra_cflags,
extra_cuda_cflags=DEFAULT_CUDA_CFLAGS + extra_cuda_cflags,
extra_cuda_cflags=selected_cuda_cflags + extra_cuda_cflags,
extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags,
extra_include_paths=DEFAULT_INCLUDE + extra_include_paths,
build_directory=build_directory,