v4.4 update. (#2979)
This commit is contained in:
@@ -123,6 +123,19 @@ function(cutlass_add_cutlass_library)
|
||||
PRIVATE cutlass_library_internal_interface
|
||||
)
|
||||
|
||||
# Use medium code model to handle large libraries (>2GB) on Linux x86_64
|
||||
# Conservative approach: only apply on Linux with GCC/Clang/Intel compilers
|
||||
# This addresses the "relocation truncated to fit" linker error when libraries exceed 2GB
|
||||
if(CMAKE_SYSTEM_NAME STREQUAL "Linux" AND
|
||||
CMAKE_SYSTEM_PROCESSOR MATCHES "x86_64|AMD64" AND
|
||||
CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang|IntelLLVM")
|
||||
message(STATUS "CUTLASS: Applying -mcmodel=medium for large library support (Linux x86_64, ${CMAKE_CXX_COMPILER_ID})")
|
||||
target_compile_options(${__NAME}_objs PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-mcmodel=medium>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:-Xcompiler=-mcmodel=medium>
|
||||
)
|
||||
endif()
|
||||
|
||||
if (CUTLASS_BUILD_MONO_LIBRARY AND __SUFFIX)
|
||||
|
||||
# If we're only building a single monolithic library then we
|
||||
|
||||
@@ -56,6 +56,10 @@ namespace library {
|
||||
// 2. e3m2_e4m3_f32_f16_e5m2
|
||||
// 3. e3m2_e4m3_f32_f16_f16
|
||||
// 4. e3m2_e4m3_f32_f32_f32
|
||||
// 5. e3m2_e5m2_f32_f16_e4m3
|
||||
// 6. e3m2_e5m2_f32_f16_e5m2
|
||||
// 7. e3m2_e5m2_f32_f16_f16
|
||||
// 8. e3m2_e5m2_f32_f32_f32
|
||||
|
||||
void initialize_gemm_reference_operations_f6_f8_f32(Manifest &manifest) {
|
||||
|
||||
@@ -138,6 +142,86 @@ void initialize_gemm_reference_operations_f6_f8_f32(Manifest &manifest) {
|
||||
float, // ElementAccumulator
|
||||
float // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 5.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
half_t, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float_e4m3_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 6.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
half_t, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float_e5m2_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 7.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
half_t, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
half_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 8.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
float, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 5.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
void, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float_e4m3_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 6.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
void, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float_e5m2_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 7.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
void, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
half_t // ElementD
|
||||
>(manifest);
|
||||
|
||||
// 8.
|
||||
make_gemm_real_canonical_layouts<
|
||||
float_e3m2_t, // ElementA
|
||||
float_e5m2_t, // ElementB
|
||||
void, // ElementC
|
||||
float, // ElementScalar
|
||||
float, // ElementAccumulator
|
||||
float // ElementD
|
||||
>(manifest);
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -318,8 +318,30 @@ protected:
|
||||
int swizzle_size,
|
||||
bool is_dynamic_cluster_enabled);
|
||||
|
||||
/// Verifies CUTLASS against host and device references
|
||||
bool verify_with_reference_(
|
||||
/// Verifies CUTLASS against device reference (for regular grouped GEMM)
|
||||
bool verify_with_device_reference_(
|
||||
Options const& options,
|
||||
PerformanceReport& report,
|
||||
DeviceContext& device_context,
|
||||
library::Operation const* operation,
|
||||
ProblemSpace const& problem_space,
|
||||
ProblemSpace::Problem const& problem,
|
||||
cutlass::library::NumericTypeID element_A,
|
||||
cutlass::library::NumericTypeID element_B);
|
||||
|
||||
/// Verifies CUTLASS against host reference (for regular grouped GEMM)
|
||||
bool verify_regular_with_host_reference_(
|
||||
Options const& options,
|
||||
PerformanceReport& report,
|
||||
DeviceContext& device_context,
|
||||
library::Operation const* operation,
|
||||
ProblemSpace const& problem_space,
|
||||
ProblemSpace::Problem const& problem,
|
||||
cutlass::library::NumericTypeID element_A,
|
||||
cutlass::library::NumericTypeID element_B);
|
||||
|
||||
/// Verifies CUTLASS against host reference (for block-scaled/blockwise grouped GEMM)
|
||||
bool verify_block_with_host_reference_(
|
||||
Options const& options,
|
||||
PerformanceReport& report,
|
||||
DeviceContext& device_context,
|
||||
|
||||
@@ -1137,37 +1137,31 @@ bool BlockScaledGemmOperationProfiler::verify_cutlass(
|
||||
}
|
||||
|
||||
|
||||
bool verification_status = verify_with_reference_(options, report, device_context, operation, problem_space, problem, element_A, element_B);
|
||||
|
||||
// Update disposition to worst case verification outcome among all
|
||||
// verification providers which are supported
|
||||
bool is_any_verification_run_passed = false;
|
||||
for (auto &m : results_.back().verification_map) {
|
||||
if (m.second == Disposition::kFailed || m.second == Disposition::kIncorrect) {
|
||||
results_.back().disposition = m.second;
|
||||
return true;
|
||||
}
|
||||
if (!is_any_verification_run_passed && m.second == Disposition::kPassed) {
|
||||
is_any_verification_run_passed = true;
|
||||
}
|
||||
if (!verify_with_reference_(options, report, device_context, operation, problem_space, problem, element_A, element_B)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (is_any_verification_run_passed) {
|
||||
results_.back().disposition = Disposition::kPassed;
|
||||
// Block-scaled GEMM only supports host verification - check it directly
|
||||
switch (results_.back().verification_map[library::Provider::kReferenceHost]) {
|
||||
case Disposition::kFailed:
|
||||
results_.back().disposition = Disposition::kFailed;
|
||||
return true;
|
||||
case Disposition::kIncorrect:
|
||||
results_.back().disposition = Disposition::kIncorrect;
|
||||
return true;
|
||||
case Disposition::kPassed:
|
||||
results_.back().disposition = Disposition::kPassed;
|
||||
return true;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// if verification.required is set, then return success iff at least one ref-check was run
|
||||
if (options.verification.required) {
|
||||
bool did_any_verification_run = false;
|
||||
for (auto provider : options.verification.providers) {
|
||||
did_any_verification_run |= (Disposition::kNotRun != results_.back().verification_map[provider]);
|
||||
}
|
||||
|
||||
if (not did_any_verification_run) {
|
||||
results_.back().status = Status::kErrorNotSupported;
|
||||
return false;
|
||||
}
|
||||
// if verification.required is set, check if host verification ran (the only supported provider)
|
||||
if (options.verification.required &&
|
||||
results_.back().verification_map[library::Provider::kReferenceHost] == Disposition::kNotRun) {
|
||||
results_.back().status = Status::kErrorNotSupported;
|
||||
return false;
|
||||
}
|
||||
|
||||
// Return true means continue profiling
|
||||
@@ -1664,3 +1658,4 @@ Status BlockScaledGemmOperationProfiler::profile_cutlass_(
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -988,37 +988,31 @@ bool BlockwiseGemmOperationProfiler::verify_cutlass(
|
||||
}
|
||||
|
||||
|
||||
bool verification_status = verify_with_reference_(options, report, device_context, operation, problem_space, problem, element_A, element_B);
|
||||
|
||||
// Update disposition to worst case verification outcome among all
|
||||
// verification providers which are supported
|
||||
bool is_any_verification_run_passed = false;
|
||||
for (auto &m : results_.back().verification_map) {
|
||||
if (m.second == Disposition::kFailed || m.second == Disposition::kIncorrect) {
|
||||
results_.back().disposition = m.second;
|
||||
return true;
|
||||
}
|
||||
if (!is_any_verification_run_passed && m.second == Disposition::kPassed) {
|
||||
is_any_verification_run_passed = true;
|
||||
}
|
||||
if (!verify_with_reference_(options, report, device_context, operation, problem_space, problem, element_A, element_B)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (is_any_verification_run_passed) {
|
||||
results_.back().disposition = Disposition::kPassed;
|
||||
// Blockwise GEMM only supports host verification - check it directly
|
||||
switch (results_.back().verification_map[library::Provider::kReferenceHost]) {
|
||||
case Disposition::kFailed:
|
||||
results_.back().disposition = Disposition::kFailed;
|
||||
return true;
|
||||
case Disposition::kIncorrect:
|
||||
results_.back().disposition = Disposition::kIncorrect;
|
||||
return true;
|
||||
case Disposition::kPassed:
|
||||
results_.back().disposition = Disposition::kPassed;
|
||||
return true;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// if verification.required is set, then return success iff at least one ref-check was run
|
||||
if (options.verification.required) {
|
||||
bool did_any_verification_run = false;
|
||||
for (auto provider : options.verification.providers) {
|
||||
did_any_verification_run |= (Disposition::kNotRun != results_.back().verification_map[provider]);
|
||||
}
|
||||
|
||||
if (not did_any_verification_run) {
|
||||
results_.back().status = Status::kErrorNotSupported;
|
||||
return false;
|
||||
}
|
||||
// if verification.required is set, check if host verification ran (the only supported provider)
|
||||
if (options.verification.required &&
|
||||
results_.back().verification_map[library::Provider::kReferenceHost] == Disposition::kNotRun) {
|
||||
results_.back().status = Status::kErrorNotSupported;
|
||||
return false;
|
||||
}
|
||||
|
||||
// Return true means continue profiling
|
||||
@@ -1555,3 +1549,4 @@ Status BlockwiseGemmOperationProfiler::profile_cutlass_(
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -35,6 +35,7 @@
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cstdint>
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/arch/mma_sm90.hpp"
|
||||
@@ -175,21 +176,79 @@ static void dequantize(DequantizedElement* dq_buffer,
|
||||
CUDA_CHECK(cudaStreamSynchronize(stream));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
template <typename ElementScale, typename QuantType = cutlass::int4b_t>
|
||||
class packed_scale_t {
|
||||
public:
|
||||
static_assert(cute::sizeof_bits_v<T> == 8,
|
||||
"only 8 bit arithmetic types are supported.");
|
||||
static_assert(
|
||||
cute::sizeof_bits_v<ElementScale> == 8,
|
||||
"ElementScale must be a supported 8-bit type.");
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
explicit packed_scale_t(T val) {
|
||||
if constexpr (!cute::is_unsigned_v<T>) {
|
||||
// Only pack negative values. The positive values are generated in flight in the mainloop.
|
||||
storage[0] = pack4(T(float(val) * -8.f), T(float(val) * -7.f), T(float(val) * -6.f), T(float(val) * -5.f));
|
||||
storage[1] = pack4(T(float(val) * -4.f), T(float(val) * -3.f), T(float(val) * -2.f), -val);
|
||||
}
|
||||
else {
|
||||
storage[0] = pack4(T(float(val) * 8.f), T(float(val) * 7.f), T(float(val) * 6.f), T(float(val) * 5.f));
|
||||
storage[1] = pack4(T(float(val) * 4.f), T(float(val) * 3.f), T(float(val) * 2.f), val);
|
||||
explicit packed_scale_t(ElementScale val) {
|
||||
|
||||
if constexpr (cutlass::platform::is_floating_point<QuantType>::value ||
|
||||
cute::is_same_v<QuantType, cutlass::float_e2m1_t>) {
|
||||
// floating point QuantType needs to be converted to a signed type
|
||||
// or a floating point ttype
|
||||
static_assert(
|
||||
!std::is_unsigned_v<ElementScale> ||
|
||||
cutlass::platform::is_floating_point<ElementScale>::value ||
|
||||
cute::is_same_v<ElementScale, cutlass::float_e4m3_t> ||
|
||||
cute::is_same_v<ElementScale, cutlass::float_e5m2_t>,
|
||||
"E2M1 quantization requires ElementScale to be signed or FP8");
|
||||
|
||||
// E2M1 quantization: Use E2M1 LUT values
|
||||
storage[0] = pack4(
|
||||
ElementScale(float(val) * 0.0f), // E2M1: 0b000 => 0.0
|
||||
ElementScale(float(val) * (-0.5f)), // E2M1: 0b001 => -0.5
|
||||
ElementScale(float(val) * (-1.0f)), // E2M1: 0b010 => -1.0
|
||||
ElementScale(float(val) * (-1.5f)) // E2M1: 0b011 => -1.5
|
||||
);
|
||||
storage[1] = pack4(
|
||||
ElementScale(float(val) * (-2.0f)), // E2M1: 0b100 => -2.0
|
||||
ElementScale(float(val) * (-3.0f)), // E2M1: 0b101 => -3.0
|
||||
ElementScale(float(val) * (-4.0f)), // E2M1: 0b110 => -4.0
|
||||
ElementScale(float(val) * (-6.0f)) // E2M1: 0b111 => -6.0
|
||||
);
|
||||
} else if constexpr (!std::is_unsigned_v<ElementScale>) {
|
||||
static_assert(
|
||||
(std::is_integral_v<QuantType> && std::is_signed_v<QuantType>) ||
|
||||
cute::is_same_v<QuantType, cutlass::int4b_t>,
|
||||
"Two's complement LUT requires signed integer QuantType (int4b_t)");
|
||||
|
||||
// INT4 Two's Complement quantization: Use TC LUT values
|
||||
storage[0] = pack4(
|
||||
ElementScale(float(val) * (-8.0f)), // TC: 0b000 => -8
|
||||
ElementScale(float(val) * (-7.0f)), // TC: 0b001 => -7
|
||||
ElementScale(float(val) * (-6.0f)), // TC: 0b010 => -6
|
||||
ElementScale(float(val) * (-5.0f)) // TC: 0b011 => -5
|
||||
);
|
||||
storage[1] = pack4(
|
||||
ElementScale(float(val) * (-4.0f)), // TC: 0b100 => -4
|
||||
ElementScale(float(val) * (-3.0f)), // TC: 0b101 => -3
|
||||
ElementScale(float(val) * (-2.0f)), // TC: 0b110 => -2
|
||||
ElementScale(float(val) * (-1.0f)) // TC: 0b111 => -1
|
||||
);
|
||||
} else {
|
||||
// Unsigned ElementScale and QuantType: Pack positive LUT values (for UINT4)
|
||||
static_assert(((std::is_integral_v<QuantType> && std::is_unsigned_v<QuantType>) ||
|
||||
cute::is_same_v<QuantType, cutlass::uint4b_t>) &&
|
||||
std::is_unsigned_v<ElementScale>,
|
||||
"Uint4 LUT requires unsigned QuantType (uint4b_t) and ElementScale");
|
||||
|
||||
// UINT4 LUT: pack positive LUT values
|
||||
storage[0] = pack4(
|
||||
ElementScale(float(val) * 8.0f),
|
||||
ElementScale(float(val) * 7.0f),
|
||||
ElementScale(float(val) * 6.0f),
|
||||
ElementScale(float(val) * 5.0f)
|
||||
);
|
||||
storage[1] = pack4(
|
||||
ElementScale(float(val) * 4.0f),
|
||||
ElementScale(float(val) * 3.0f),
|
||||
ElementScale(float(val) * 2.0f),
|
||||
ElementScale(float(val) * 1.0f)
|
||||
);
|
||||
}
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -230,7 +289,7 @@ private:
|
||||
Storage storage[2] {};
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Storage pack4(T c1, T c2, T c3, T c4) {
|
||||
static Storage pack4(ElementScale c1, ElementScale c2, ElementScale c3, ElementScale c4) {
|
||||
Storage result = 0;
|
||||
result |= (static_cast<Storage>(reinterpret_cast<Stage const&>(c4)) << 24);
|
||||
result |= (static_cast<Storage>(reinterpret_cast<Stage const&>(c3)) << 16);
|
||||
@@ -239,25 +298,25 @@ private:
|
||||
return result;
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
T get() const {
|
||||
ElementScale get() const {
|
||||
auto stage = static_cast<Stage>(storage[0] >> 8);
|
||||
#if defined(__CUDA_ARCH__)
|
||||
return reinterpret_cast<T const&>(stage);
|
||||
return reinterpret_cast<ElementScale const&>(stage);
|
||||
#else
|
||||
T tmp;
|
||||
ElementScale tmp;
|
||||
std::memcpy(&tmp, &stage, sizeof(Stage));
|
||||
return tmp;
|
||||
#endif
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
T get(int idx) const {
|
||||
ElementScale get(int idx) const {
|
||||
Stage stage;
|
||||
if (idx < 4) stage = static_cast<Stage>(storage[0] >> (8 * idx));
|
||||
else stage = static_cast<Stage>(storage[1] >> (8 * idx - 32));
|
||||
#if defined(__CUDA_ARCH__)
|
||||
return reinterpret_cast<T const&>(stage);
|
||||
return reinterpret_cast<ElementScale const&>(stage);
|
||||
#else
|
||||
T tmp;
|
||||
ElementScale tmp;
|
||||
std::memcpy(&tmp, &stage, sizeof(Stage));
|
||||
return tmp;
|
||||
#endif
|
||||
@@ -301,7 +360,7 @@ static bool unified_encode_int4b(cutlass::int4b_t const *block_in, cutlass::int4
|
||||
return true;
|
||||
}
|
||||
|
||||
template <class ElementScale>
|
||||
template <class ElementScale, typename QuantType = cutlass::int4b_t>
|
||||
static bool pack_scale_fp8(ElementScale const *block_in, cutlass::Array<ElementScale, 8> *block_out, const size_t block_size) {
|
||||
std::vector<ElementScale> data_in(block_size);
|
||||
std::vector<cutlass::Array<ElementScale, 8>> data_out(block_size);
|
||||
@@ -315,7 +374,7 @@ static bool pack_scale_fp8(ElementScale const *block_in, cutlass::Array<ElementS
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < block_size; i++) {
|
||||
cutlass::packed_scale_t<ElementScale> tmp(data_in[i]);
|
||||
cutlass::packed_scale_t<ElementScale, QuantType> tmp(data_in[i]);
|
||||
data_out[i] = reinterpret_cast<cutlass::Array<ElementScale, 8> const&>(tmp);
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user