Skip to content

Commit 32a0db4

Browse files
committed
workaround only for sm120 and 121
1 parent f083600 commit 32a0db4

2 files changed

Lines changed: 42 additions & 13 deletions

File tree

cpp/include/raft/linalg/detail/cublaslt_wrappers.hpp

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
#include <raft/core/resource/cublaslt_handle.hpp>
1111
#include <raft/core/resource/cuda_stream.hpp>
1212
#include <raft/core/resource/custom_resource.hpp>
13+
#include <raft/core/resource/device_properties.hpp>
1314
#include <raft/core/resources.hpp>
1415
#include <raft/util/cache.hpp>
1516
#include <raft/util/cuda_data_type.hpp>
@@ -104,12 +105,16 @@ struct matmul_key_hash {
104105
* cuBLASLt 13.6 may select algorithm 68 once A's physical span reaches 2^31 elements. That
105106
* algorithm fails during execution for FP32, so select the next ranked heuristic instead.
106107
*/
107-
inline auto needs_cublaslt_13_6_workaround(const matmul_key_t& args, std::size_t version) noexcept
108-
-> bool
108+
inline auto needs_cublaslt_13_6_workaround(const matmul_key_t& args,
109+
std::size_t version,
110+
int device_major,
111+
int device_minor) noexcept -> bool
109112
{
110113
constexpr uint64_t max_safe_span = (uint64_t{1} << 31) - 1;
111114
const auto a_columns = args.trans_a ? args.m : args.k;
112-
return version == 130600 && args.lda != 0 && a_columns > max_safe_span / args.lda;
115+
const bool is_sm120_or_sm121 = device_major == 12 && (device_minor == 0 || device_minor == 1);
116+
return version == 130600 && is_sm120_or_sm121 && args.lda != 0 &&
117+
a_columns > max_safe_span / args.lda;
113118
}
114119

115120
inline auto get_cublaslt_algorithm_id(const cublasLtMatmulHeuristicResult_t& heuristic) -> int
@@ -212,7 +217,9 @@ struct matmul_desc {
212217
bool use_cublaslt_13_6_workaround = false;
213218
if constexpr (std::is_same_v<S, float> && std::is_same_v<A, float> &&
214219
std::is_same_v<B, float> && std::is_same_v<C, float>) {
215-
use_cublaslt_13_6_workaround = needs_cublaslt_13_6_workaround(args, cublasLtGetVersion());
220+
const auto& device_properties = resource::get_device_properties(res);
221+
use_cublaslt_13_6_workaround = needs_cublaslt_13_6_workaround(
222+
args, cublasLtGetVersion(), device_properties.major, device_properties.minor);
216223
}
217224

218225
constexpr int workaround_heuristic_results = 2;

cpp/tests/linalg/gemm_basic.cpp

Lines changed: 31 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -166,31 +166,53 @@ TEST(Raft, GemmPointerModeDeviceDefaults) { test_gemm_pointer_mode_device(false,
166166
TEST(Raft, GemmCublasLt136WorkaroundPredicate)
167167
{
168168
constexpr std::size_t affected_version = 130600;
169+
constexpr int affected_device_major = 12;
170+
constexpr int affected_device_minor = 1;
169171
const detail::matmul_key_t below_boundary{134217727, 1, 2, 16, 1, 134217727, true, true};
170172
const detail::matmul_key_t at_boundary{134217728, 1, 2, 16, 1, 134217728, true, true};
171173
const detail::matmul_key_t above_boundary{134217729, 1, 2, 16, 1, 134217729, true, true};
172174

173-
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(below_boundary, affected_version));
174-
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version));
175-
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(above_boundary, affected_version));
176-
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, 130599));
177-
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, 130601));
175+
const auto needs_workaround = [&](const auto& args) {
176+
return detail::needs_cublaslt_13_6_workaround(
177+
args, affected_version, affected_device_major, affected_device_minor);
178+
};
179+
180+
EXPECT_FALSE(needs_workaround(below_boundary));
181+
EXPECT_TRUE(needs_workaround(at_boundary));
182+
EXPECT_TRUE(needs_workaround(above_boundary));
183+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(
184+
at_boundary, 130599, affected_device_major, affected_device_minor));
185+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(
186+
at_boundary, 130601, affected_device_major, affected_device_minor));
178187

179188
auto different_output = at_boundary;
180189
different_output.trans_b = false;
181190
different_output.n = 2;
182191
different_output.ldb = 7;
183192
different_output.ldc = 11;
184-
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(different_output, affected_version));
193+
EXPECT_TRUE(needs_workaround(different_output));
185194

186195
const detail::matmul_key_t non_transposed_below{2, 1, 134217727, 16, 134217727, 2, false, false};
187196
const detail::matmul_key_t non_transposed_at{2, 1, 134217728, 16, 134217728, 2, false, false};
188-
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(non_transposed_below, affected_version));
189-
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(non_transposed_at, affected_version));
197+
EXPECT_FALSE(needs_workaround(non_transposed_below));
198+
EXPECT_TRUE(needs_workaround(non_transposed_at));
190199

191200
auto invalid_lda = at_boundary;
192201
invalid_lda.lda = 0;
193-
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(invalid_lda, affected_version));
202+
EXPECT_FALSE(needs_workaround(invalid_lda));
203+
}
204+
205+
TEST(Raft, GemmCublasLt136WorkaroundArchitectures)
206+
{
207+
constexpr std::size_t affected_version = 130600;
208+
const detail::matmul_key_t at_boundary{134217728, 1, 2, 16, 1, 134217728, true, true};
209+
210+
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 12, 0));
211+
EXPECT_TRUE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 12, 1));
212+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 7, 5));
213+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 8, 0));
214+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 8, 9));
215+
EXPECT_FALSE(detail::needs_cublaslt_13_6_workaround(at_boundary, affected_version, 9, 0));
194216
}
195217

196218
} // namespace raft::linalg

0 commit comments

Comments
 (0)