@@ -166,31 +166,53 @@ TEST(Raft, GemmPointerModeDeviceDefaults) { test_gemm_pointer_mode_device(false,
166166TEST (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