diff --git a/python/sglang/jit_kernel/include/sgl_kernel/runtime.cuh b/python/sglang/jit_kernel/include/sgl_kernel/runtime.cuh index 4ea722a3f..36bc00dc5 100644 --- a/python/sglang/jit_kernel/include/sgl_kernel/runtime.cuh +++ b/python/sglang/jit_kernel/include/sgl_kernel/runtime.cuh @@ -26,6 +26,9 @@ #ifndef cudaDevAttrComputeCapabilityMajor #define cudaDevAttrComputeCapabilityMajor hipDeviceAttributeComputeCapabilityMajor #endif +#ifndef cudaDevAttrComputeCapabilityMinor +#define cudaDevAttrComputeCapabilityMinor hipDeviceAttributeComputeCapabilityMinor +#endif #ifndef cudaRuntimeGetVersion #define cudaRuntimeGetVersion hipRuntimeGetVersion #endif @@ -68,6 +71,18 @@ inline auto get_cc_major(int device_id) -> int { return cc_major; } +// Return the Minor compute capability for the given device +inline auto get_cc_minor(int device_id) -> int { + int cc_minor; + RuntimeDeviceCheck(cudaDeviceGetAttribute(&cc_minor, cudaDevAttrComputeCapabilityMinor, device_id)); + return cc_minor; +} + +// Return the SM version (major * 10 + minor) for the given device +inline auto get_sm_version(int device_id) -> int { + return get_cc_major(device_id) * 10 + get_cc_minor(device_id); +} + // Return the runtime version inline auto get_runtime_version() -> int { int runtime_version; diff --git a/python/sglang/jit_kernel/include/sgl_kernel/utils.cuh b/python/sglang/jit_kernel/include/sgl_kernel/utils.cuh index 82974a99a..684475117 100644 --- a/python/sglang/jit_kernel/include/sgl_kernel/utils.cuh +++ b/python/sglang/jit_kernel/include/sgl_kernel/utils.cuh @@ -242,14 +242,6 @@ inline void RuntimeDeviceCheck(DebugInfo location = {}) { return RuntimeDeviceCheck(::cudaGetLastError(), location); } -inline int getSMVersion(int device_id) { - int sm_major = 0; - int sm_minor = 0; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&sm_major, cudaDevAttrComputeCapabilityMajor, device_id)); - RuntimeDeviceCheck(cudaDeviceGetAttribute(&sm_minor, cudaDevAttrComputeCapabilityMinor, device_id)); - return sm_major * 10 + sm_minor; -} - inline auto alloc_workspace_tensor(size_t required_bytes, DLDevice device) -> tvm::ffi::Tensor { if (required_bytes == 0) return {}; DLDataType u8 = {kDLUInt, 8, 1};