Commit a1808b5
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
- kernels
- funcs/blas
- impl
- test/cpp/phi/kernels
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
153 | 153 | | |
154 | 154 | | |
155 | 155 | | |
| 156 | + | |
| 157 | + | |
| 158 | + | |
| 159 | + | |
156 | 160 | | |
157 | | - | |
| 161 | + | |
| 162 | + | |
| 163 | + | |
| 164 | + | |
| 165 | + | |
| 166 | + | |
| 167 | + | |
| 168 | + | |
158 | 169 | | |
159 | 170 | | |
160 | 171 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
229 | 229 | | |
230 | 230 | | |
231 | 231 | | |
232 | | - | |
233 | | - | |
234 | | - | |
235 | | - | |
236 | | - | |
237 | | - | |
| 232 | + | |
| 233 | + | |
238 | 234 | | |
239 | 235 | | |
240 | 236 | | |
| |||
357 | 353 | | |
358 | 354 | | |
359 | 355 | | |
360 | | - | |
361 | | - | |
362 | | - | |
| 356 | + | |
| 357 | + | |
| 358 | + | |
363 | 359 | | |
364 | 360 | | |
365 | 361 | | |
366 | 362 | | |
367 | 363 | | |
368 | | - | |
| 364 | + | |
369 | 365 | | |
370 | 366 | | |
371 | 367 | | |
| |||
0 commit comments