Skip to content

Commit a1808b5

Browse files
committed
Add 64-bit strided BatchGEMM dispatch
Fix FP16 scalar type and wrong specialization in batched GEMM cuBLAS reads alpha/beta with the width implied by computeType. In the BatchedGEMM overload templated on a separate scalar type U, the CUBLAS_COMPUTE_16F path passed the address of the U-typed alpha (float for baddbmm), so cuBLAS read the low 16 bits of a float as a half: alpha = 1.0f became half(0.0) and the result silently degenerated to beta * C. Pass fp16 copies of the scalars instead. The sibling overload was already correct because its alpha is of the element type T. The 64-bit branch of the same overload called CUBlas<phi::bfloat16>::GEMM_STRIDED_BATCH_64 from inside the T == phi::float16 branch. This is currently harmless since both specializations forward to detail::GemmStridedBatchedEx64 and the dtype comes from the CUDA_R_16F argument, but it breaks as soon as either specialization gains its own logic. Use CUBlas<T> as the sibling overload does. Also gate FP16 accumulation on compute capability 70, matching PyTorch's prop->major >= 7 check in aten/src/ATen/cuda/CUDABlas.cpp, and apply it to both overloads so the same flag does not behave differently between them. Test Plan: not covered by the existing CPU-only tests in test/cpp/phi/kernels/test_math_function.cc. The scalar-type fix is observable by running bmm/baddbmm in float16 with FLAGS_gemm_use_half_precision_compute_type=1 and comparing against torch, which currently returns beta * C. A GPU differential test for this is tracked separately. Remove unreachable ld range checks in batched GEMM lda/ldb/ldc are each derived from M/N/K under a packed layout, so the to_blas_int calls on them could only ever fire after the M/N/K calls on the preceding lines already had. Drop the dead checks. Also record why the N == 1 fast path passes the wrong ldc; fixing it changes kernel selection on a hot path and needs a benchmark first. Drop the N == 1 fast path in strided batched GEMM The branch passed the swapped-layout ldc (== N == 1) where cuBLAS wants ldc >= max(1, M) for the column-major M x 1 result, so its own guard ldc >= max(1, M) degenerated into M <= 1. It therefore never fired for the batched GEMV shapes it was written to accelerate, only for 1x1 results, and PyTorch has no such special case in bgemm at all. Removing it is behaviour-preserving and cannot regress a shape the branch could actually reach. Reintroducing it as a real GEMV optimisation means fixing ldc to M, which changes kernel selection on a hot path and belongs in a separate, benchmarked change. Add 64-bit dispatch for pointer-array BatchedGEMM/TRSM and fix HIP fp16 scalars Three defects remained after a85b469 added the 64-bit strided path: 1. The four pointer-array BatchedGEMM specializations (double, float, float16, bfloat16) and TRSM had no 64-bit path at all: they narrowed M/N/K/batchCount unconditionally, so any dimension above INT_MAX failed even though cuBLAS provides cublas{S,D}gemmBatched_64, cublasGemmBatchedEx_64 and cublas{S,D,C,Z}trsm_v2_64. Worse, they derived lda/ldb/ldc from the *already narrowed* int values, so the leading dimensions were computed in 32-bit arithmetic. They now validate non-negativity, compute leading dimensions in int64_t, and dispatch to the _64 APIs above INT_MAX. GETRF/GETRI/MatInv/GETRS keep narrowing and throwing because cuBLAS ships no _64 variants for them. 2. Both fp16 GEMM specializations in blas_impl.hip.h handed rocBLAS a float* alpha/beta while declaring rocblas_datatype_f16_r as the compute type. rocBLAS reads the scalars with the width implied by the compute type, so with FLAGS_gemm_use_half_precision_compute_type=1 it read the low 16 bits of a float: alpha=1.0f became half(0.0) and the result degenerated to beta*C. Same defect class as the CUDA fix in ee8c8a0. 3. The two strided specializations taking float alpha with fp16/bf16 inputs never checked M/N/K/batchCount for negativity, unlike every sibling specialization, letting negative extents reach cuBLAS. Also normalizes the guard idiom (a named requires_64_bit_blas bool), collapses the seven divergent Unimplemented messages to five that all name Linux (matching the defined(__linux__) half of the guard), and declares the eight newly used _64 symbols in cublas.h. Test Plan: - clang-format 21.1.7 restricted to the changed line ranges with DerivePointerAlignment:false, PointerAlignment:Right (the repo pins 13.0.0; 21 derives Left on these files and would contradict Paddle's `void *a`): zero diffs on all three files. - cpplint 2.0.2 with the repo filter from .pre-commit-config.yaml: clean. - Verified all 21 _64 symbols are used <=> declared in cublas.h <=> present in /usr/local/cuda/include/cublas_api.h <=> exported by libcublas.so.13, with parameter types checked against the prototypes. - Verified #if/#endif and brace/paren balance in both headers. - Not built: no CUDA or HIP compilation was attempted in this environment, so compile-correctness is unverified. This change was authored with AI assistance.
1 parent 4306edd commit a1808b5

7 files changed

Lines changed: 1341 additions & 517 deletions

File tree

‎paddle/phi/backends/dynload/cublas.h‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -153,8 +153,19 @@ CUBLAS_WORKSPACE_ROUTINE(DECLARE_DYNAMIC_LOAD_CUBLAS_WRAP)
153153
__macro(cublasDgemm_v2_64); \
154154
__macro(cublasCgemm_v2_64); \
155155
__macro(cublasZgemm_v2_64); \
156+
__macro(cublasSgemmStridedBatched_64); \
157+
__macro(cublasDgemmStridedBatched_64); \
158+
__macro(cublasCgemmStridedBatched_64); \
159+
__macro(cublasZgemmStridedBatched_64); \
156160
__macro(cublasGemmStridedBatchedEx_64); \
157-
__macro(cublasGemmEx_64);
161+
__macro(cublasGemmEx_64); \
162+
__macro(cublasSgemmBatched_64); \
163+
__macro(cublasDgemmBatched_64); \
164+
__macro(cublasGemmBatchedEx_64); \
165+
__macro(cublasStrsm_v2_64); \
166+
__macro(cublasDtrsm_v2_64); \
167+
__macro(cublasCtrsm_v2_64); \
168+
__macro(cublasZtrsm_v2_64);
158169

159170
CUBLAS_BLAS_ROUTINE_EACH_R5(DECLARE_DYNAMIC_LOAD_CUBLAS_WRAP)
160171
#endif

‎paddle/phi/kernels/funcs/blas/blas.h‎

Lines changed: 6 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -229,12 +229,8 @@ class Blas {
229229
#endif
230230

231231
template <typename T>
232-
void MatMul(const int M,
233-
const int N,
234-
const int K,
235-
const T* A,
236-
const T* B,
237-
T* C) const;
232+
void MatMul(
233+
int64_t M, int64_t N, int64_t K, const T* A, const T* B, T* C) const;
238234

239235
template <typename T>
240236
void MatMul(const DenseTensor& mat_a,
@@ -357,15 +353,15 @@ class Blas {
357353
template <typename T>
358354
void BatchedGEMM(CBLAS_TRANSPOSE transA,
359355
CBLAS_TRANSPOSE transB,
360-
int M,
361-
int N,
362-
int K,
356+
int64_t M,
357+
int64_t N,
358+
int64_t K,
363359
T alpha,
364360
const T** A,
365361
const T** B,
366362
T beta,
367363
T** C,
368-
int batchCount) const;
364+
int64_t batchCount) const;
369365

370366
#if defined(PADDLE_WITH_MKLML) && !defined(PADDLE_WITH_CUDA) && \
371367
!defined(PADDLE_WITH_HIP)

0 commit comments

Comments
 (0)