kernels: scalar scale A support for fp8_gemm (#33469)
Signed-off-by: Devashish Lal <devcode@fb.com> Co-authored-by: Devashish Lal <devcode@fb.com>
This commit is contained in:
co-authored by
Devashish Lal
parent
269d51ed4b
commit
fc74c35546
@@ -8,9 +8,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from sgl_kernel import fp8_scaled_mm as sgl_scaled_mm
|
from sgl_kernel import fp8_scaled_mm as sgl_scaled_mm
|
||||||
|
|
||||||
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
|
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import per_tensor_quant_fp8
|
||||||
per_tensor_quant_fp8,
|
|
||||||
)
|
|
||||||
from sglang.utils import is_in_ci
|
from sglang.utils import is_in_ci
|
||||||
|
|
||||||
# Optional vLLM import
|
# Optional vLLM import
|
||||||
@@ -106,7 +104,7 @@ def sglang_scaled_fp8_quant(
|
|||||||
if IS_CI:
|
if IS_CI:
|
||||||
batch_sizes = [1] # Single batch size for CI
|
batch_sizes = [1] # Single batch size for CI
|
||||||
else:
|
else:
|
||||||
batch_sizes = [1, 16, 64, 128, 256, 512, 1024, 2048]
|
batch_sizes = [1, 2, 8, 16, 64, 128, 256, 512, 1024, 2048]
|
||||||
|
|
||||||
# Filter line_vals based on vLLM availability
|
# Filter line_vals based on vLLM availability
|
||||||
if VLLM_AVAILABLE:
|
if VLLM_AVAILABLE:
|
||||||
@@ -115,24 +113,39 @@ if VLLM_AVAILABLE:
|
|||||||
"vllm-fp8-bf16",
|
"vllm-fp8-bf16",
|
||||||
"sglang-fp8-fp16",
|
"sglang-fp8-fp16",
|
||||||
"sglang-fp8-bf16",
|
"sglang-fp8-bf16",
|
||||||
|
"sglang-scalar-a-fp8-fp16",
|
||||||
|
"sglang-scalar-a-fp8-bf16",
|
||||||
]
|
]
|
||||||
line_names = [
|
line_names = [
|
||||||
"vllm-fp8-fp16",
|
"vllm-fp8-fp16",
|
||||||
"vllm-fp8-bf16",
|
"vllm-fp8-bf16",
|
||||||
"sglang-fp8-fp16",
|
"sglang-fp8-fp16",
|
||||||
"sglang-fp8-bf16",
|
"sglang-fp8-bf16",
|
||||||
|
"sglang-scalar-a-fp8-fp16",
|
||||||
|
"sglang-scalar-a-fp8-bf16",
|
||||||
|
]
|
||||||
|
styles = [
|
||||||
|
("green", "-"),
|
||||||
|
("green", "--"),
|
||||||
|
("blue", "-"),
|
||||||
|
("blue", "--"),
|
||||||
|
("red", "-"),
|
||||||
|
("red", "--"),
|
||||||
]
|
]
|
||||||
styles = [("green", "-"), ("green", "--"), ("blue", "-"), ("blue", "--")]
|
|
||||||
else:
|
else:
|
||||||
line_vals = [
|
line_vals = [
|
||||||
"sglang-fp8-fp16",
|
"sglang-fp8-fp16",
|
||||||
"sglang-fp8-bf16",
|
"sglang-fp8-bf16",
|
||||||
|
"sglang-scalar-a-fp8-fp16",
|
||||||
|
"sglang-scalar-a-fp8-bf16",
|
||||||
]
|
]
|
||||||
line_names = [
|
line_names = [
|
||||||
"sglang-fp8-fp16",
|
"sglang-fp8-fp16",
|
||||||
"sglang-fp8-bf16",
|
"sglang-fp8-bf16",
|
||||||
|
"sglang-scalar-a-fp8-fp16",
|
||||||
|
"sglang-scalar-a-fp8-bf16",
|
||||||
]
|
]
|
||||||
styles = [("blue", "-"), ("blue", "--")]
|
styles = [("blue", "-"), ("blue", "--"), ("red", "-"), ("red", "--")]
|
||||||
|
|
||||||
|
|
||||||
@triton.testing.perf_report(
|
@triton.testing.perf_report(
|
||||||
@@ -174,8 +187,9 @@ def benchmark(batch_size, provider, N, K):
|
|||||||
lambda: vllm_scaled_mm(a_fp8, b_fp8, scale_a_fp8, scale_b_fp8, dtype),
|
lambda: vllm_scaled_mm(a_fp8, b_fp8, scale_a_fp8, scale_b_fp8, dtype),
|
||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
elif "sglang-fp8" in provider:
|
elif "sglang" in provider:
|
||||||
a_fp8, scale_a_fp8 = sglang_scaled_fp8_quant(a, scale_a)
|
a_scale = scale_a_scalar if "scalar-a" in provider else scale_a
|
||||||
|
a_fp8, scale_a_fp8 = sglang_scaled_fp8_quant(a, a_scale)
|
||||||
b_fp8, scale_b_fp8 = sglang_scaled_fp8_quant(b, scale_b)
|
b_fp8, scale_b_fp8 = sglang_scaled_fp8_quant(b, scale_b)
|
||||||
b_fp8 = b_fp8.t()
|
b_fp8 = b_fp8.t()
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||||
|
|||||||
@@ -448,19 +448,22 @@ template <
|
|||||||
typename MainloopScheduleType,
|
typename MainloopScheduleType,
|
||||||
typename EpilogueScheduleType,
|
typename EpilogueScheduleType,
|
||||||
typename TileSchedulerType = void,
|
typename TileSchedulerType = void,
|
||||||
bool WithBias = false>
|
bool WithBias = false,
|
||||||
|
bool ScalarA = false>
|
||||||
struct DeviceGemmFp8RowwiseSm100 {
|
struct DeviceGemmFp8RowwiseSm100 {
|
||||||
static_assert(std::is_same_v<ElementType, cutlass::float_e4m3_t>, "ElementType must be FP8(e4m3)");
|
static_assert(std::is_same_v<ElementType, cutlass::float_e4m3_t>, "ElementType must be FP8(e4m3)");
|
||||||
using TileShape = CTAShape;
|
using TileShape = CTAShape;
|
||||||
using Accum = cutlass::epilogue::fusion::Sm90AccFetch;
|
using Accum = cutlass::epilogue::fusion::Sm90AccFetch;
|
||||||
|
|
||||||
using ElementComputeEpilogue = float;
|
using ElementComputeEpilogue = float;
|
||||||
using ScaleA = cutlass::epilogue::fusion::Sm90ColBroadcast<
|
using VectorScaleA = cutlass::epilogue::fusion::Sm90ColBroadcast<
|
||||||
0,
|
0,
|
||||||
TileShape,
|
TileShape,
|
||||||
ElementComputeEpilogue,
|
ElementComputeEpilogue,
|
||||||
ElementComputeEpilogue,
|
ElementComputeEpilogue,
|
||||||
cute::Stride<cute::Int<1>, cute::Int<0>, cute::Int<0>>>;
|
cute::Stride<cute::Int<1>, cute::Int<0>, cute::Int<0>>>;
|
||||||
|
using ScalarScaleA = cutlass::epilogue::fusion::Sm90ScalarBroadcast<float>;
|
||||||
|
using ScaleA = std::conditional_t<ScalarA, ScalarScaleA, VectorScaleA>;
|
||||||
|
|
||||||
using ScaleB = cutlass::epilogue::fusion::Sm90RowBroadcast<
|
using ScaleB = cutlass::epilogue::fusion::Sm90RowBroadcast<
|
||||||
0,
|
0,
|
||||||
@@ -551,7 +554,11 @@ struct DeviceGemmFp8RowwiseSm100 {
|
|||||||
auto* data_ptr = static_cast<T*>(tensor.data_ptr());
|
auto* data_ptr = static_cast<T*>(tensor.data_ptr());
|
||||||
static_assert(
|
static_assert(
|
||||||
std::is_same_v<Descriptor, ScaleA> || std::is_same_v<Descriptor, ScaleB> || std::is_same_v<Descriptor, Bias>);
|
std::is_same_v<Descriptor, ScaleA> || std::is_same_v<Descriptor, ScaleB> || std::is_same_v<Descriptor, Bias>);
|
||||||
return Arguments{data_ptr};
|
if constexpr (std::is_same_v<Descriptor, ScalarScaleA>) {
|
||||||
|
return Arguments{{}, {data_ptr}, {}};
|
||||||
|
} else {
|
||||||
|
return Arguments{data_ptr};
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
public:
|
public:
|
||||||
@@ -657,7 +664,7 @@ void launch_sm100_fp8_scaled_mm(
|
|||||||
TORCH_CHECK(status == cutlass::Status::kSuccess)
|
TORCH_CHECK(status == cutlass::Status::kSuccess)
|
||||||
}
|
}
|
||||||
|
|
||||||
template <typename OutType>
|
template <typename OutType, bool ScalarA>
|
||||||
void sm100_fp8_dispatch_bias(
|
void sm100_fp8_dispatch_bias(
|
||||||
torch::Tensor& out,
|
torch::Tensor& out,
|
||||||
const torch::Tensor& a,
|
const torch::Tensor& a,
|
||||||
@@ -695,7 +702,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
true>;
|
true,
|
||||||
|
ScalarA>;
|
||||||
using BiasGemm256 = DeviceGemmFp8RowwiseSm100<
|
using BiasGemm256 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -705,7 +713,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
true>;
|
true,
|
||||||
|
ScalarA>;
|
||||||
using BiasGemm64 = DeviceGemmFp8RowwiseSm100<
|
using BiasGemm64 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -715,7 +724,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
true>;
|
true,
|
||||||
|
ScalarA>;
|
||||||
using BiasGemm16 = DeviceGemmFp8RowwiseSm100<
|
using BiasGemm16 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -725,7 +735,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
true>;
|
true,
|
||||||
|
ScalarA>;
|
||||||
|
|
||||||
// Gemm type without bias
|
// Gemm type without bias
|
||||||
using GemmDefault = DeviceGemmFp8RowwiseSm100<
|
using GemmDefault = DeviceGemmFp8RowwiseSm100<
|
||||||
@@ -737,7 +748,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
false>;
|
false,
|
||||||
|
ScalarA>;
|
||||||
using Gemm256 = DeviceGemmFp8RowwiseSm100<
|
using Gemm256 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -747,7 +759,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
false>;
|
false,
|
||||||
|
ScalarA>;
|
||||||
using Gemm64 = DeviceGemmFp8RowwiseSm100<
|
using Gemm64 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -757,7 +770,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
false>;
|
false,
|
||||||
|
ScalarA>;
|
||||||
using Gemm16 = DeviceGemmFp8RowwiseSm100<
|
using Gemm16 = DeviceGemmFp8RowwiseSm100<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
ElementOutput,
|
ElementOutput,
|
||||||
@@ -767,7 +781,8 @@ void sm100_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
false>;
|
false,
|
||||||
|
ScalarA>;
|
||||||
|
|
||||||
// next power of 2 (minimum 16)
|
// next power of 2 (minimum 16)
|
||||||
uint32_t const m = a.size(0);
|
uint32_t const m = a.size(0);
|
||||||
@@ -811,7 +826,10 @@ void sm100_fp8_dispatch_shape(
|
|||||||
const torch::Tensor& scales_a,
|
const torch::Tensor& scales_a,
|
||||||
const torch::Tensor& scales_b,
|
const torch::Tensor& scales_b,
|
||||||
const c10::optional<torch::Tensor>& bias) {
|
const c10::optional<torch::Tensor>& bias) {
|
||||||
return sm100_fp8_dispatch_bias<OutType>(out, a, b, scales_a, scales_b, bias);
|
if (scales_a.numel() == 1) {
|
||||||
|
return sm100_fp8_dispatch_bias<OutType, true>(out, a, b, scales_a, scales_b, bias);
|
||||||
|
}
|
||||||
|
return sm100_fp8_dispatch_bias<OutType, false>(out, a, b, scales_a, scales_b, bias);
|
||||||
}
|
}
|
||||||
|
|
||||||
template <
|
template <
|
||||||
@@ -823,19 +841,22 @@ template <
|
|||||||
typename MainloopScheduleType,
|
typename MainloopScheduleType,
|
||||||
typename EpilogueScheduleType,
|
typename EpilogueScheduleType,
|
||||||
typename TileSchedulerType = void,
|
typename TileSchedulerType = void,
|
||||||
bool WithBias = false>
|
bool WithBias = false,
|
||||||
|
bool ScalarA = false>
|
||||||
struct DeviceGemmFp8RowwiseSm120 {
|
struct DeviceGemmFp8RowwiseSm120 {
|
||||||
static_assert(std::is_same_v<ElementType, cutlass::float_e4m3_t>, "ElementType must be FP8(e4m3)");
|
static_assert(std::is_same_v<ElementType, cutlass::float_e4m3_t>, "ElementType must be FP8(e4m3)");
|
||||||
using TileShape = CTAShape;
|
using TileShape = CTAShape;
|
||||||
using Accum = cutlass::epilogue::fusion::Sm90AccFetch;
|
using Accum = cutlass::epilogue::fusion::Sm90AccFetch;
|
||||||
|
|
||||||
using ElementComputeEpilogue = float;
|
using ElementComputeEpilogue = float;
|
||||||
using ScaleA = cutlass::epilogue::fusion::Sm90ColBroadcast<
|
using VectorScaleA = cutlass::epilogue::fusion::Sm90ColBroadcast<
|
||||||
0,
|
0,
|
||||||
TileShape,
|
TileShape,
|
||||||
ElementComputeEpilogue,
|
ElementComputeEpilogue,
|
||||||
ElementComputeEpilogue,
|
ElementComputeEpilogue,
|
||||||
cute::Stride<cute::Int<1>, cute::Int<0>, cute::Int<0>>>;
|
cute::Stride<cute::Int<1>, cute::Int<0>, cute::Int<0>>>;
|
||||||
|
using ScalarScaleA = cutlass::epilogue::fusion::Sm90ScalarBroadcast<float>;
|
||||||
|
using ScaleA = std::conditional_t<ScalarA, ScalarScaleA, VectorScaleA>;
|
||||||
|
|
||||||
using ScaleB = cutlass::epilogue::fusion::Sm90RowBroadcast<
|
using ScaleB = cutlass::epilogue::fusion::Sm90RowBroadcast<
|
||||||
0,
|
0,
|
||||||
@@ -926,7 +947,11 @@ struct DeviceGemmFp8RowwiseSm120 {
|
|||||||
auto* data_ptr = static_cast<T*>(tensor.data_ptr());
|
auto* data_ptr = static_cast<T*>(tensor.data_ptr());
|
||||||
static_assert(
|
static_assert(
|
||||||
std::is_same_v<Descriptor, ScaleA> || std::is_same_v<Descriptor, ScaleB> || std::is_same_v<Descriptor, Bias>);
|
std::is_same_v<Descriptor, ScaleA> || std::is_same_v<Descriptor, ScaleB> || std::is_same_v<Descriptor, Bias>);
|
||||||
return Arguments{data_ptr};
|
if constexpr (std::is_same_v<Descriptor, ScalarScaleA>) {
|
||||||
|
return Arguments{{}, {data_ptr}, {}};
|
||||||
|
} else {
|
||||||
|
return Arguments{data_ptr};
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
public:
|
public:
|
||||||
@@ -1032,7 +1057,7 @@ void launch_sm120_fp8_scaled_mm(
|
|||||||
TORCH_CHECK(status == cutlass::Status::kSuccess)
|
TORCH_CHECK(status == cutlass::Status::kSuccess)
|
||||||
}
|
}
|
||||||
|
|
||||||
template <typename OutType>
|
template <typename OutType, bool ScalarA>
|
||||||
void sm120_fp8_dispatch_bias(
|
void sm120_fp8_dispatch_bias(
|
||||||
torch::Tensor& out,
|
torch::Tensor& out,
|
||||||
const torch::Tensor& a,
|
const torch::Tensor& a,
|
||||||
@@ -1060,7 +1085,8 @@ void sm120_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
true>;
|
true,
|
||||||
|
ScalarA>;
|
||||||
|
|
||||||
using GemmDefault = DeviceGemmFp8RowwiseSm120<
|
using GemmDefault = DeviceGemmFp8RowwiseSm120<
|
||||||
ElementInput,
|
ElementInput,
|
||||||
@@ -1071,7 +1097,8 @@ void sm120_fp8_dispatch_bias(
|
|||||||
MainloopScheduleType,
|
MainloopScheduleType,
|
||||||
EpilogueScheduleType,
|
EpilogueScheduleType,
|
||||||
TileSchedulerType,
|
TileSchedulerType,
|
||||||
false>;
|
false,
|
||||||
|
ScalarA>;
|
||||||
|
|
||||||
if (bias) {
|
if (bias) {
|
||||||
return launch_sm120_fp8_scaled_mm<BiasGemmDefault, true>(out, a, b, scales_a, scales_b, bias);
|
return launch_sm120_fp8_scaled_mm<BiasGemmDefault, true>(out, a, b, scales_a, scales_b, bias);
|
||||||
@@ -1088,7 +1115,10 @@ void sm120_fp8_dispatch_shape(
|
|||||||
const torch::Tensor& scales_a,
|
const torch::Tensor& scales_a,
|
||||||
const torch::Tensor& scales_b,
|
const torch::Tensor& scales_b,
|
||||||
const c10::optional<torch::Tensor>& bias) {
|
const c10::optional<torch::Tensor>& bias) {
|
||||||
return sm120_fp8_dispatch_bias<OutType>(out, a, b, scales_a, scales_b, bias);
|
if (scales_a.numel() == 1) {
|
||||||
|
return sm120_fp8_dispatch_bias<OutType, true>(out, a, b, scales_a, scales_b, bias);
|
||||||
|
}
|
||||||
|
return sm120_fp8_dispatch_bias<OutType, false>(out, a, b, scales_a, scales_b, bias);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
@@ -1115,7 +1145,26 @@ torch::Tensor fp8_scaled_mm(
|
|||||||
TORCH_CHECK(mat_b.scalar_type() == torch::kFloat8_e4m3fn, "mat_b must be Float8_e4m3fn");
|
TORCH_CHECK(mat_b.scalar_type() == torch::kFloat8_e4m3fn, "mat_b must be Float8_e4m3fn");
|
||||||
TORCH_CHECK(out_dtype == torch::kHalf || out_dtype == torch::kBFloat16, "out_dtype must be Half or BFloat16");
|
TORCH_CHECK(out_dtype == torch::kHalf || out_dtype == torch::kBFloat16, "out_dtype must be Half or BFloat16");
|
||||||
|
|
||||||
TORCH_CHECK(scales_a.numel() == mat_a.size(0), "size of scales_a is not matched");
|
auto sm_version = getSMVersion();
|
||||||
|
TORCH_CHECK(
|
||||||
|
scales_a.numel() == 1 || scales_a.numel() == mat_a.size(0),
|
||||||
|
"scales_a must contain either one scalar scale or one scale per row; got ",
|
||||||
|
scales_a.numel(),
|
||||||
|
" elements for M=",
|
||||||
|
mat_a.size(0));
|
||||||
|
bool scalar_a_scale_supported = false;
|
||||||
|
#if defined CUDA_VERSION && CUDA_VERSION >= 12000
|
||||||
|
scalar_a_scale_supported = sm_version == 90;
|
||||||
|
#endif
|
||||||
|
#if defined CUDA_VERSION && CUDA_VERSION >= 12080
|
||||||
|
scalar_a_scale_supported = scalar_a_scale_supported || sm_version >= 100;
|
||||||
|
#endif
|
||||||
|
TORCH_CHECK(
|
||||||
|
scales_a.numel() != 1 || mat_a.size(0) == 1 || scalar_a_scale_supported,
|
||||||
|
"scalar scales_a with M > 1 is unsupported on SM",
|
||||||
|
sm_version,
|
||||||
|
" for this build; got M=",
|
||||||
|
mat_a.size(0));
|
||||||
TORCH_CHECK(scales_b.numel() == mat_b.size(1), "size of scales_b is not matched");
|
TORCH_CHECK(scales_b.numel() == mat_b.size(1), "size of scales_b is not matched");
|
||||||
TORCH_CHECK(scales_a.is_contiguous(), "scales_a must be contiguous");
|
TORCH_CHECK(scales_a.is_contiguous(), "scales_a must be contiguous");
|
||||||
TORCH_CHECK(scales_b.is_contiguous(), "scales_b msut be contiguous");
|
TORCH_CHECK(scales_b.is_contiguous(), "scales_b msut be contiguous");
|
||||||
@@ -1131,8 +1180,6 @@ torch::Tensor fp8_scaled_mm(
|
|||||||
torch::Tensor out = torch::empty({mat_a.size(0), mat_b.size(1)}, mat_a.options().dtype(out_dtype));
|
torch::Tensor out = torch::empty({mat_a.size(0), mat_b.size(1)}, mat_a.options().dtype(out_dtype));
|
||||||
TORCH_CHECK((out.size(1) * out.element_size()) % 16 == 0, "out must be multiple of 16 bytes for memory alignment");
|
TORCH_CHECK((out.size(1) * out.element_size()) % 16 == 0, "out must be multiple of 16 bytes for memory alignment");
|
||||||
|
|
||||||
auto sm_version = getSMVersion();
|
|
||||||
|
|
||||||
#if defined CUDA_VERSION && CUDA_VERSION >= 12080
|
#if defined CUDA_VERSION && CUDA_VERSION >= 12080
|
||||||
if (sm_version >= 120) {
|
if (sm_version >= 120) {
|
||||||
if (out_dtype == torch::kBFloat16) {
|
if (out_dtype == torch::kBFloat16) {
|
||||||
|
|||||||
@@ -5,6 +5,22 @@ import torch
|
|||||||
from sgl_kernel import fp8_scaled_mm
|
from sgl_kernel import fp8_scaled_mm
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_version_at_least(major, minor):
|
||||||
|
if torch.version.cuda is None:
|
||||||
|
return False
|
||||||
|
version = tuple(int(component) for component in torch.version.cuda.split(".")[:2])
|
||||||
|
return version >= (major, minor)
|
||||||
|
|
||||||
|
|
||||||
|
def _native_scalar_a_supported():
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
return False
|
||||||
|
capability = torch.cuda.get_device_capability()
|
||||||
|
if capability == (9, 0):
|
||||||
|
return _cuda_version_at_least(12, 0)
|
||||||
|
return capability[0] in (10, 12) and _cuda_version_at_least(12, 8)
|
||||||
|
|
||||||
|
|
||||||
def torch_scaled_mm(a, b, scale_a, scale_b, out_dtype, bias):
|
def torch_scaled_mm(a, b, scale_a, scale_b, out_dtype, bias):
|
||||||
o = torch.matmul(a.to(torch.float32), b.to(torch.float32))
|
o = torch.matmul(a.to(torch.float32), b.to(torch.float32))
|
||||||
o = o.to(torch.float32)
|
o = o.to(torch.float32)
|
||||||
@@ -38,6 +54,40 @@ def _test_accuracy_once(M, N, K, with_bias, out_dtype, device):
|
|||||||
print(f"M={M}, N={N}, K={K}, with_bias={with_bias}, out_dtype={out_dtype}: OK")
|
print(f"M={M}, N={N}, K={K}, with_bias={with_bias}, out_dtype={out_dtype}: OK")
|
||||||
|
|
||||||
|
|
||||||
|
def _test_scalar_a_accuracy_once(M, N, K, with_bias, out_dtype, device):
|
||||||
|
fp8_info = torch.finfo(torch.float8_e4m3fn)
|
||||||
|
a_fp8 = (
|
||||||
|
torch.randn(M, K, dtype=torch.float32, device=device)
|
||||||
|
.clamp(min=fp8_info.min, max=fp8_info.max)
|
||||||
|
.to(torch.float8_e4m3fn)
|
||||||
|
)
|
||||||
|
b_fp8 = (
|
||||||
|
torch.randn(N, K, dtype=torch.float32, device=device)
|
||||||
|
.clamp(min=fp8_info.min, max=fp8_info.max)
|
||||||
|
.to(torch.float8_e4m3fn)
|
||||||
|
.t()
|
||||||
|
)
|
||||||
|
scale_a = torch.tensor([0.03125], device=device, dtype=torch.float32)
|
||||||
|
scale_a_repeated = scale_a.repeat(M)
|
||||||
|
|
||||||
|
# Resemble merged projections whose component matrices were quantized with
|
||||||
|
# different tensorwise scales before concatenation.
|
||||||
|
scale_b = torch.empty(N, device=device, dtype=torch.float32)
|
||||||
|
first_boundary = N // 3
|
||||||
|
second_boundary = 2 * N // 3
|
||||||
|
scale_b[:first_boundary] = 0.015625
|
||||||
|
scale_b[first_boundary:second_boundary] = 0.03125
|
||||||
|
scale_b[second_boundary:] = 0.0625
|
||||||
|
|
||||||
|
bias = torch.randn(N, device=device, dtype=out_dtype) if with_bias else None
|
||||||
|
expected = torch_scaled_mm(a_fp8, b_fp8, scale_a, scale_b, out_dtype, bias)
|
||||||
|
actual = fp8_scaled_mm(a_fp8, b_fp8, scale_a, scale_b, out_dtype, bias)
|
||||||
|
repeated = fp8_scaled_mm(a_fp8, b_fp8, scale_a_repeated, scale_b, out_dtype, bias)
|
||||||
|
|
||||||
|
torch.testing.assert_close(expected, actual, rtol=0.02, atol=1)
|
||||||
|
torch.testing.assert_close(repeated, actual, rtol=0, atol=0)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("M", [1, 128, 512, 1024, 4096])
|
@pytest.mark.parametrize("M", [1, 128, 512, 1024, 4096])
|
||||||
@pytest.mark.parametrize("N", [16, 128, 512, 1024, 4096])
|
@pytest.mark.parametrize("N", [16, 128, 512, 1024, 4096])
|
||||||
@pytest.mark.parametrize("K", [512, 1024, 4096, 8192, 16384])
|
@pytest.mark.parametrize("K", [512, 1024, 4096, 8192, 16384])
|
||||||
@@ -92,6 +142,43 @@ def test_accuracy_sm90_swap_ab(shape_mn, K, with_bias, out_dtype):
|
|||||||
_test_accuracy_once(M, N, K, with_bias, out_dtype, "cuda")
|
_test_accuracy_once(M, N, K, with_bias, out_dtype, "cuda")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
not _native_scalar_a_supported(),
|
||||||
|
reason="native scalar A scales require a compatible SM90, SM100, or SM120 build",
|
||||||
|
)
|
||||||
|
@pytest.mark.parametrize("M", [1, 2, 8, 16, 64, 189])
|
||||||
|
@pytest.mark.parametrize("with_bias", [True, False])
|
||||||
|
@pytest.mark.parametrize("out_dtype", [torch.bfloat16, torch.float16])
|
||||||
|
def test_scalar_a_channelwise_b(M, with_bias, out_dtype):
|
||||||
|
_test_scalar_a_accuracy_once(M, 6144, 4096, with_bias, out_dtype, "cuda")
|
||||||
|
|
||||||
|
|
||||||
|
def test_rejects_invalid_a_scale_count():
|
||||||
|
M, N, K = 8, 128, 512
|
||||||
|
a = torch.randn(M, K, device="cuda").to(torch.float8_e4m3fn)
|
||||||
|
b = torch.randn(N, K, device="cuda").to(torch.float8_e4m3fn).t()
|
||||||
|
scale_a = torch.ones(2, device="cuda", dtype=torch.float32)
|
||||||
|
scale_b = torch.ones(N, device="cuda", dtype=torch.float32)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="scales_a must contain either"):
|
||||||
|
fp8_scaled_mm(a, b, scale_a, scale_b, torch.bfloat16, None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (8, 9),
|
||||||
|
reason="SM89-specific scalar A validation",
|
||||||
|
)
|
||||||
|
def test_rejects_scalar_a_with_multiple_rows_on_sm89():
|
||||||
|
M, N, K = 8, 128, 512
|
||||||
|
a = torch.randn(M, K, device="cuda").to(torch.float8_e4m3fn)
|
||||||
|
b = torch.randn(N, K, device="cuda").to(torch.float8_e4m3fn).t()
|
||||||
|
scale_a = torch.ones(1, device="cuda", dtype=torch.float32)
|
||||||
|
scale_b = torch.ones(N, device="cuda", dtype=torch.float32)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="scalar scales_a with M > 1 is unsupported"):
|
||||||
|
fp8_scaled_mm(a, b, scale_a, scale_b, torch.bfloat16, None)
|
||||||
|
|
||||||
|
|
||||||
PRODUCTION_LIKE_FP8_GEMM_CASES = [
|
PRODUCTION_LIKE_FP8_GEMM_CASES = [
|
||||||
(189, 4608, 8192, False, torch.bfloat16),
|
(189, 4608, 8192, False, torch.bfloat16),
|
||||||
(3330, 256, 8192, False, torch.bfloat16),
|
(3330, 256, 8192, False, torch.bfloat16),
|
||||||
|
|||||||
Reference in New Issue
Block a user