Thanks for using Compiler Explorer
Sponsors
Jakt
C++
Ada
Algol68
Analysis
Android Java
Android Kotlin
Assembly
C
C3
Carbon
C with Coccinelle
C++ with Coccinelle
C++ (Circle)
CIRCT
Clean
Clojure
CO2
Cargo
CMake
CMakeScript
COBOL
C++ for OpenCL
Makefile
Maven
MLIR
Cppx
Cppx-Blue
Cppx-Gold
Cpp2-cppfront
Crystal
C#
CUDA C++
CuTe DSL
D
Dart
Elixir
Erlang
Fortran
F#
GLSL
Go
Haskell
HLSL
Helion
Hook
Hylo
IL
ispc
Java
Julia
Kotlin
Lean
LLVM IR
LLVM MIR
Lua
Modula-2
Mojo
Nim
Numba
Nix
Objective-C
Objective-C++
OCaml
Odin
OpenCL C
Pascal
Perl
Pony
PTX
Python
Racket
Raku
RazorForge
Ruby
Rust
Sail
SFPI C++
Snowball
Scala
Slang
Solidity
Spice
SPIR-V
Swift
LLVM TableGen
Toit
Triton
TypeScript Native
V
Vala
Visual Basic
Vyper
WASM
Yul (Solidity IR)
Zig
Javascript
GIMPLE
Ygen
sway
cuda source #1
Output
Compile to binary object
Link to binary
Execute the code
Intel asm syntax
Demangle identifiers
Verbose demangling
Filters
Unused labels
Library functions
Directives
Comments
Horizontal whitespace
Debug intrinsics
Compiler
10.0.0 sm_75 CUDA-10.2
10.0.1 sm_75 CUDA-10.2
11.0.0 sm_75 CUDA-10.2
16.0.0 sm_90 CUDA-11.8
17.0.1(libc++) sm_90 CUDA-12.1
18.1.0(libc++) sm_90 CUDA-12.3.1
19.1.0 sm_90 CUDA-12.5.1
20.1.0 sm_90 CUDA-12.5.1
20.1.0 sm_90 CUDA-12.6.1
20.1.0 sm_90 CUDA-12.6.2
20.1.0 sm_90 CUDA-12.8.1
20.1.0 sm_90 CUDA-12.9.0
20.1.0 sm_90 CUDA-12.9.1
NVCC 10.0.130
NVCC 10.1.105
NVCC 10.1.168
NVCC 10.1.243
NVCC 10.2.89
NVCC 11.0.2
NVCC 11.0.3
NVCC 11.1.0
NVCC 11.1.1
NVCC 11.2.0
NVCC 11.2.1
NVCC 11.2.2
NVCC 11.3.0
NVCC 11.3.1
NVCC 11.4.0
NVCC 11.4.1
NVCC 11.4.2
NVCC 11.4.3
NVCC 11.4.4
NVCC 11.5.0
NVCC 11.5.1
NVCC 11.5.2
NVCC 11.6.0
NVCC 11.6.1
NVCC 11.6.2
NVCC 11.7.0
NVCC 11.7.1
NVCC 11.8.0
NVCC 12.0.0
NVCC 12.0.1
NVCC 12.1.0
NVCC 12.2.1
NVCC 12.3.1
NVCC 12.4.1
NVCC 12.5.1
NVCC 12.6.1
NVCC 12.6.2
NVCC 12.8.1
NVCC 12.9.0
NVCC 12.9.1
NVCC 13.0.0
NVCC 13.0.1
NVCC 13.0.2
NVCC 13.1.0
NVCC 13.1.1
NVCC 13.2.0
NVCC 13.3.0
NVCC 9.1.85
NVCC 9.2.88
NVRTC 11.0.2
NVRTC 11.0.3
NVRTC 11.1.0
NVRTC 11.1.1
NVRTC 11.2.0
NVRTC 11.2.1
NVRTC 11.2.2
NVRTC 11.3.0
NVRTC 11.3.1
NVRTC 11.4.0
NVRTC 11.4.1
NVRTC 11.5.0
NVRTC 11.5.1
NVRTC 11.5.2
NVRTC 11.6.0
NVRTC 11.6.1
NVRTC 11.6.2
NVRTC 11.7.0
NVRTC 11.7.1
NVRTC 11.8.0
NVRTC 12.0.0
NVRTC 12.0.1
NVRTC 12.1.0
NVRTC 12.2.1
NVRTC 12.3.1
NVRTC 12.4.1
NVRTC 12.5.1
NVRTC 12.6.1
NVRTC 12.6.2
NVRTC 12.8.1
NVRTC 12.9.0
NVRTC 12.9.1
NVRTC 13.0.0
NVRTC 13.0.1
NVRTC 13.0.2
NVRTC 13.1.0
NVRTC 13.1.1
NVRTC 13.2.0
NVRTC 13.3.0
SCALE NVCC (AMD) 1.7.1
SCALE NVCC (AMD) 1.7.2
SCALE NVCC (NVIDIA) 1.7.1
SCALE NVCC (NVIDIA) 1.7.2
clang 7.0.0 sm_70 CUDA-9.1
clang 8.0.0 sm_75 CUDA-10.0
clang 9.0.0 sm_75 CUDA-10.1
clang rocm-10.0.0
clang rocm-4.5.2
clang rocm-5.0.2
clang rocm-5.1.3
clang rocm-5.2.3
clang rocm-5.3.2
clang rocm-5.7.0
clang rocm-6.0.2
clang rocm-6.1.2
clang rocm-6.2.4
clang rocm-6.3.3
clang rocm-6.4.0
clang rocm-7.0.1
clang rocm-7.0.2
clang rocm-7.1.0
clang rocm-7.1.1
clang rocm-7.14.0
clang rocm-7.2.0
clang rocm-7.2.1
clang staging rocm-10.0.0
clang staging rocm-6.1.2
clang staging rocm-6.2.4
clang staging rocm-6.3.3
clang staging rocm-6.4.0
clang staging rocm-7.0.1
clang staging rocm-7.0.2
clang staging rocm-7.1.0
clang staging rocm-7.1.1
clang staging rocm-7.14.0
clang staging rocm-7.2.0
clang staging rocm-7.2.1
clang trunk rocm-10.0.0
clang trunk rocm-6.1.2
clang trunk rocm-6.2.4
clang trunk rocm-6.3.3
clang trunk rocm-6.4.0
clang trunk rocm-7.0.1
clang trunk rocm-7.0.2
clang trunk rocm-7.1.0
clang trunk rocm-7.1.1
clang trunk rocm-7.14.0
clang trunk rocm-7.2.0
clang trunk rocm-7.2.1
trunk sm_120a CUDA-13.0.0
trunk sm_120a CUDA-13.0.1
trunk sm_120a CUDA-13.0.2
trunk sm_120a CUDA-13.1.0
Options
Source code
#define CEIL_DIV(value, divisor) (((value) + (divisor) - 1) / (divisor)) __global__ void sgemm_vectorised(const float *__restrict__ A, const float *__restrict__ B, float *__restrict__ C, int M, int N, int K, float alpha, float beta) { const uint TILE_SIZE_N = 8; const uint ROWS_PER_THREAD = 8; const uint COLS_PER_THREAD = 8; const uint TILE_SIZE_M = 128; const uint TILE_SIZE_K = 128; // Allocate shared memory __shared__ float sharedA[TILE_SIZE_M * TILE_SIZE_N]; __shared__ float sharedB[TILE_SIZE_N * TILE_SIZE_K]; // Identify the tile of C this thread block is responsible for const uint block_row = blockIdx.y; const uint block_column = blockIdx.x; // Calculate position of thread within tile (Remapping from 1-D to 2-D) Note --> Each thread is a grid in itself hanlding ROWS_PER_THREAD x COLS_PER_THREAD const uint ty = threadIdx.x / (TILE_SIZE_K / COLS_PER_THREAD); // 0, ..., 15 const uint tx = threadIdx.x % (TILE_SIZE_K / COLS_PER_THREAD); // 0, ..., 15 // Move pointers from A[0], B[0] and C[0] to the starting positions of the tile A += block_row * TILE_SIZE_M * N; // Move pointer (block_row * TILE_SIZE_M) rows down B += block_column * TILE_SIZE_K; // Move pointer (block_column * TILE_SIZE_K) columns to the right C += (block_row * TILE_SIZE_M * K) + (block_column * TILE_SIZE_K); // Move pointer (block_row * TILE_SIZE_M * K) rows down then (block_column * TILE_SIZE_K) columns to the right // Map each thread to one 4-float chunk that it will load. const uint smem_ty_A = threadIdx.x / (TILE_SIZE_N / 4); // --> 0, ..., 127 const uint smem_tx_A = threadIdx.x % (TILE_SIZE_N / 4); // --> 0, 1 const uint smem_ty_B = threadIdx.x / (TILE_SIZE_K / 4); // --> 0, ..., 7 const uint smem_tx_B = threadIdx.x % (TILE_SIZE_K / 4); // --> 0, ..., 31 // Calculate how many tiles we have const uint num_tiles = CEIL_DIV(N, TILE_SIZE_N); float thread_results[ROWS_PER_THREAD * COLS_PER_THREAD] = {0.0f}; float reg_m[ROWS_PER_THREAD] = {0.0f}; float reg_k[COLS_PER_THREAD] = {0.0f}; // Outer loop iterate over tiles for (int t = 0; t < num_tiles; t++) { // Populate smem using vector loads float4 tempA = reinterpret_cast<const float4 *>(&A[smem_ty_A * N + smem_tx_A * 4])[0]; // [0] dereference issues one ld.global.nc.v4.f32 // Transpose A (instead of 128x8 previously for ex, now it will be 8x128) sharedA[(smem_tx_A * 4 + 0) * TILE_SIZE_M + smem_ty_A] = tempA.x; sharedA[(smem_tx_A * 4 + 1) * TILE_SIZE_M + smem_ty_A] = tempA.y; sharedA[(smem_tx_A * 4 + 2) * TILE_SIZE_M + smem_ty_A] = tempA.z; sharedA[(smem_tx_A * 4 + 3) * TILE_SIZE_M + smem_ty_A] = tempA.w; float4 tempB = reinterpret_cast<const float4 *>(&B[smem_ty_B * K + smem_tx_B * 4])[0]; reinterpret_cast<float4 *>(&sharedB[smem_ty_B * TILE_SIZE_K + smem_tx_B * 4])[0] = tempB; __syncthreads(); // Outer loop over shared dimension N for (int i = 0; i < TILE_SIZE_N; i++) { // Load into registers one "col" (its acc row now) from sharedA and one row from sharedB // We can actually also use vectorised loads here to cut smem load instructions by 4x. for (int row = 0; row < ROWS_PER_THREAD; row += 4) { uint global_smem_row_idx = ty * ROWS_PER_THREAD + row; // i will be the same for the whole "column" although since its transposed we are accessing same row. // Notice how we also skip rows by TILE_SIZE_M now float4 temp_shared_A = reinterpret_cast<float4 *>(&sharedA[i * TILE_SIZE_M + global_smem_row_idx])[0]; // ld.shared.v4.f32 reg_m[row + 0] = temp_shared_A.x; reg_m[row + 1] = temp_shared_A.y; reg_m[row + 2] = temp_shared_A.z; reg_m[row + 3] = temp_shared_A.w; } for (int col = 0; col < COLS_PER_THREAD; col += 4) { // We can do same vectorised loads uint global_smem_col_idx = tx * COLS_PER_THREAD + col; float4 temp_shared_B = reinterpret_cast<float4 *>(&sharedB[i * TILE_SIZE_K + global_smem_col_idx])[0]; reg_k[col + 0] = temp_shared_B.x; reg_k[col + 1] = temp_shared_B.y; reg_k[col + 2] = temp_shared_B.z; reg_k[col + 3] = temp_shared_B.w; } // Calculate outer product between reg_m and reg_k to produce the partial results matrix of the thread for (uint m = 0; m < ROWS_PER_THREAD; m++) { for (uint k = 0; k < COLS_PER_THREAD; k++) { thread_results[m * COLS_PER_THREAD + k] += reg_m[m] * reg_k[k]; // --> (ROWS_PER_THREAD x COLS_PER_THREAD) matrix } } } __syncthreads(); A += TILE_SIZE_N; // Move right B += TILE_SIZE_N * K; // Move down } // Write results of the thread back to C for (uint row = 0; row < ROWS_PER_THREAD; row++) { // handle COLS_PER_THREAD in chunks of 4 for (uint col = 0; col < COLS_PER_THREAD; col += 4) { uint global_row_idx = ty * ROWS_PER_THREAD + row; uint global_col_idx = tx * COLS_PER_THREAD + col; float4 tempC = reinterpret_cast<float4 *>(&C[global_row_idx * K + global_col_idx])[0]; tempC.x = (alpha * thread_results[row * COLS_PER_THREAD + col]) + (beta * tempC.x); tempC.y = (alpha * thread_results[row * COLS_PER_THREAD + col + 1]) + (beta * tempC.y); tempC.z = (alpha * thread_results[row * COLS_PER_THREAD + col + 2]) + (beta * tempC.z); tempC.w = (alpha * thread_results[row * COLS_PER_THREAD + col + 3]) + (beta * tempC.w); reinterpret_cast<float4 *>(&C[global_row_idx * K + global_col_idx])[0] = tempC; } } }
Become a Patron
Sponsor on GitHub
Donate via PayPal
Compiler Explorer Shop
Source on GitHub
Mailing list
Installed libraries
Wiki
Report an issue
How it works
Contact the author
CE on Mastodon
CE on Bluesky
Statistics
Changelog
Version tree