Since GPUArrays v11.5.11 (PR #760), mul! with a strided (non-contiguous) view of a CuArray runs GPUArrays' generic gpu_coalesced_matmul_kernel instead of cuBLAS. cuBLAS handles these views natively: CUDA.jl's StridedCuMatrix methods pass the leading dimension, and calling cuBLAS.gemm! on the same view works. For Float32 the generic kernel is about 10× slower here.
I am not sure this should be fixed in the CUDA side or GPUArrays side. But I guess this affects all device backends so I opened an issue here.
MWE
using CUDA, LinearAlgebra
A = CUDA.rand(Float32, 4096, 1024)
Av = view(A, 1:4000, :) # strided (non-contiguous) view of a CuMatrix
Ac = copy(Av) # contiguous copy of the same data
B = CUDA.rand(Float32, 1024, 1024)
C = CUDA.zeros(Float32, 4000, 1024)
display(CUDA.@profile mul!(C, Av, B)) # view
display(CUDA.@profile mul!(C, Ac, B)) # contiguous copy
display(CUDA.@profile CUDA.cuBLAS.gemm!('N', 'N', 1f0, Av, B, 0f0, C)) # cuBLAS on the view itself
Device-side kernel (NVIDIA RTX A6000):
| call |
GPUArrays v11.5.8 |
GPUArrays v11.5.14 |
mul!(C, Av, B), view |
ampere_sgemm_128x64_nn, 0.40 ms |
gpu_coalesced_matmul_kernel, 3.98 ms |
mul!(C, Ac, B), contiguous copy |
ampere_sgemm_128x64_nn, 0.40 ms |
ampere_sgemm_128x64_nn, 0.40 ms |
cuBLAS.gemm! on Av |
ampere_sgemm_128x64_nn, 0.40 ms |
ampere_sgemm_128x64_nn, 0.40 ms |
Since GPUArrays v11.5.11 (PR #760),
mul!with a strided (non-contiguous) view of aCuArrayruns GPUArrays' genericgpu_coalesced_matmul_kernelinstead of cuBLAS. cuBLAS handles these views natively: CUDA.jl'sStridedCuMatrixmethods pass the leading dimension, and callingcuBLAS.gemm!on the same view works. ForFloat32the generic kernel is about 10× slower here.I am not sure this should be fixed in the CUDA side or GPUArrays side. But I guess this affects all device backends so I opened an issue here.
MWE
Device-side kernel (NVIDIA RTX A6000):
mul!(C, Av, B), viewampere_sgemm_128x64_nn, 0.40 msgpu_coalesced_matmul_kernel, 3.98 msmul!(C, Ac, B), contiguous copyampere_sgemm_128x64_nn, 0.40 msampere_sgemm_128x64_nn, 0.40 mscuBLAS.gemm!onAvampere_sgemm_128x64_nn, 0.40 msampere_sgemm_128x64_nn, 0.40 ms