code/day48-fusion/fusion.cuThis is the source used by the lesson and its recorded evidence. Compile commands and expected output live in the directory README.
// SPDX-License-Identifier: MIT
//
// Day 48: kernel fusion and launch overhead, priced apart.
//
// One four-stage elementwise chain, written twice:
//
// staged four kernels. Each reads its input from global memory and
// writes its output back, so the chain moves 36 bytes per
// element and pays four launches.
// fused one kernel. The three intermediates never leave a register, so
// the chain moves 12 bytes per element and pays one launch.
//
// Fusing deletes two different things, and they are not the same size. This
// program prices them separately.
//
// Part 1 what a launch costs on its own, from a kernel that does nothing.
// Part 2 the chain at 16 Mi elements, where the traffic is milliseconds
// and the launches are microseconds.
// Part 3 the same two chains swept down in size until that reverses.
// Part 4 what fusion costs. The fused kernel is widened until each thread
// holds enough live values to change how many warps fit on an SM.
//
// The byte counts are the prediction, so they are computed from the
// constants and printed above the times. Stage 3 reads the residual as well
// as its own input, so the staged chain moves 9 floats per element and not
// 8: a model that takes one stage and multiplies by four gets 32 bytes and
// is wrong before it ever meets a GPU.
//
// Build: nvcc -std=c++17 -O3 -arch=sm_75 -o fusion fusion.cu
// Run: ./fusion
//
#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <vector>
#include <cuda_runtime.h>
// The one error macro. This file is standalone, the way a Compiler Explorer
// embed is, so it carries its own verbatim copy. `err_` carries a trailing
// underscore so it cannot collide with a variable at the call site, and the
// do/while makes the macro one statement so it survives a braceless `if`.
#define CUDA_CHECK(call) \
do { \
cudaError_t err_ = (call); \
if (err_ != cudaSuccess) { \
std::fprintf(stderr, "CUDA error %s:%d: %s: %s\n", __FILE__, \
__LINE__, #call, cudaGetErrorString(err_)); \
std::exit(EXIT_FAILURE); \
} \
} while (0)
// 16 Mi elements is 64 MiB per buffer and six buffers on the device. It is
// chosen large enough that the chain's traffic is milliseconds while its
// launches stay microseconds, which is the gap part 2 is about.
constexpr size_t kElems = 16777216;
// The correctness size. 611 is 13 x 47, so no block size and no elements-
// per-thread width divides it and every bounds check runs on every launch.
// The timed size above is a power of two on purpose, so no guard fires
// inside a timed kernel and the timings are not measuring tail handling.
constexpr size_t kCheckElems = 1048576 + 611;
constexpr int kThreadsPerBlock = 256; // 8 warps
constexpr int kWarmupRuns = 3;
constexpr int kTimedRuns = 10;
// Part 1 launches an empty kernel this many times inside one timed call, so
// the event pair brackets something far larger than its own half-microsecond
// resolution. timeKernel then repeats that batch, which is why the total is
// (kWarmupRuns + kTimedRuns) x kLaunchesPerBatch launches.
constexpr int kLaunchesPerBatch = 1000;
// Part 3. Every value must fit the buffers allocated for kElems.
constexpr size_t kSweep[] = {1024, 16384, 262144, 4194304, 16777216};
constexpr int kSweepCount =
static_cast<int>(sizeof(kSweep) / sizeof(kSweep[0]));
// Part 4. A template argument cannot come from a loop, so these six widths
// are instantiated by hand in main and this array exists only so that a
// static_assert can hold the two lists together.
constexpr int kWideK[] = {1, 2, 4, 8, 16, 32};
constexpr int kWideCount = static_cast<int>(sizeof(kWideK) / sizeof(kWideK[0]));
// The chain's parameters. kHi is 3 rather than something safely above the
// range so that the clamp actually fires on about a quarter of the elements;
// a stage that never does anything is a stage the compiler may delete, and
// then the byte model is describing a program that no longer exists.
constexpr float kGamma = 1.5f;
constexpr float kBeta = -0.25f;
constexpr float kLo = -6.0f;
constexpr float kHi = 3.0f;
// The tanh form of GELU, which is what GPT-2 used and what PyTorch calls
// approximate='tanh'. kGeluA is sqrt(2 / pi) rounded to float; the CPU
// reference promotes these same rounded floats to double rather than using
// the exact constants, so the reference is checking the kernel's arithmetic
// and not the author's choice of literal.
constexpr float kGeluA = 0.7978845608f;
constexpr float kGeluB = 0.044715f;
// Floats moved per element. Staged: 2 for stage 1, 2 for stage 2, 3 for
// stage 3 because it also reads the residual, 2 for stage 4. Fused: x, the
// residual, and the output. These two numbers are the whole prediction.
constexpr double kStagedFloatsPerElem = 9.0;
constexpr double kFusedFloatsPerElem = 3.0;
constexpr double kCopyFloatsPerElem = 2.0;
constexpr double kRelTolerance = 1e-5;
constexpr double kAbsTolerance = 1e-6;
static_assert(kThreadsPerBlock % 32 == 0,
"the block is a whole number of warps");
static_assert(kThreadsPerBlock <= 1024,
"1024 threads is the per-block maximum on every compute "
"capability this course targets");
static_assert(kElems % 32 == 0,
"the timed size must be a whole number of blocks at every "
"block size this file invites, down to one warp, so that no "
"bounds check runs inside a timed kernel");
static_assert(kCheckElems % 2 == 1,
"the correctness size must be odd, so that no power-of-two "
"launch shape covers it exactly and every bounds check fires "
"on every launch");
static_assert(kSweep[kSweepCount - 1] <= kElems,
"the sweep reuses the buffers allocated for kElems, so its "
"largest size cannot exceed them");
static_assert(kWideK[0] == 1 && kWideK[1] == 2 && kWideK[2] == 4 &&
kWideK[3] == 8 && kWideK[4] == 16 && kWideK[5] == 32,
"the six explicit template instantiations in main must match "
"this list, because a template argument cannot come from a "
"loop and nothing else keeps the two in step");
static_assert(kHi > kLo, "the clamp range must be a range");
// The four stages, as ordinary device functions, so that the staged kernels
// and every fused kernel below run byte-identical arithmetic and the only
// thing that differs between them is where the intermediates live.
__device__ __forceinline__ float scaleBias(float v) {
return v * kGamma + kBeta;
}
__device__ __forceinline__ float gelu(float v) {
const float inner = kGeluA * (v + kGeluB * v * v * v);
return 0.5f * v * (1.0f + tanhf(inner));
}
__device__ __forceinline__ float addResidual(float v, float r) {
return v + r;
}
__device__ __forceinline__ float clampRange(float v) {
return fminf(fmaxf(v, kLo), kHi);
}
// Does nothing, and is the point of part 1.
//
// Memory: none. It issues no loads and no stores, so whatever the clock says
// about a run of these is the cost of getting a kernel started and finished.
//
// Launch assumption: none. Any grid, any block.
__global__ void doNothing() {}
// One read and one write per element, no arithmetic.
//
// Memory: consecutive lanes take consecutive elements, so one warp's 32
// addresses cover 128 contiguous bytes and cost four 32-byte sectors.
//
// The floor row. No elementwise chain over this buffer can beat a kernel
// that only reads it and writes it, and measuring the floor here rather than
// borrowing day 11's number keeps this program's block shape, buffer size
// and clock state out of the comparison.
//
// Launch assumption: any 1D grid that covers n.
__global__ void copyFloor(const float* __restrict__ in, float* __restrict__ out,
size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
out[i] = in[i];
}
}
// Stage 1 of 4. One thread scales and shifts one element.
//
// Memory: 32 consecutive lanes read 32 consecutive floats and write 32
// consecutive floats, so both sides coalesce into four sectors per warp. The
// access pattern is not what this program is measuring; the number of times
// the pattern happens is.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyScaleBias(const float* __restrict__ in,
float* __restrict__ out, size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
out[i] = scaleBias(in[i]);
}
}
// Stage 2 of 4. One thread takes the GELU of one element.
//
// Memory: as stage 1, one coalesced read and one coalesced write. This is
// the only stage with a transcendental in it, and it is here so that the
// chain has some real arithmetic to be memory-bound in spite of.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyGelu(const float* __restrict__ in, float* __restrict__ out,
size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
out[i] = gelu(in[i]);
}
}
// Stage 3 of 4. One thread adds one residual element to one input element.
//
// Memory: two coalesced reads and one coalesced write, so this stage moves
// three floats per element where the others move two. That asymmetry is
// deliberate. A staged chain of four two-float stages would move 32 bytes
// per element and a reader could get the right ratio from the wrong model.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyAddResidual(const float* __restrict__ in,
const float* __restrict__ r,
float* __restrict__ out, size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
out[i] = addResidual(in[i], r[i]);
}
}
// Stage 4 of 4. One thread clamps one element into [kLo, kHi].
//
// Memory: one coalesced read and one coalesced write.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
__global__ void applyClampRange(const float* __restrict__ in,
float* __restrict__ out, size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
out[i] = clampRange(in[i]);
}
}
// All four stages in one thread, with the three intermediates living in
// registers instead of in global memory.
//
// Memory: two coalesced reads and one coalesced write per element, 12 bytes
// against the staged chain's 36. The arithmetic is the same four device
// functions in the same order, so anything the clock finds between this and
// the staged chain is traffic and launches, not maths.
//
// Launch assumption: gridDim.x * blockDim.x >= n.
// snippet: fused-chain
__global__ void chainFused(const float* __restrict__ x,
const float* __restrict__ r, float* __restrict__ out,
size_t n) {
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
if (i < n) {
float v = scaleBias(x[i]);
v = gelu(v);
v = addResidual(v, r[i]);
out[i] = clampRange(v);
}
}
// end snippet: fused-chain
// The same fused chain, with each thread carrying kWide elements instead of
// one. Same total work, same total traffic, same arithmetic, same order.
// The only thing that changes is how many values are live at once.
//
// Memory: for a fixed k the 32 lanes of a warp read 32 consecutive floats,
// because the elements a thread owns are one whole grid apart rather than
// adjacent. That is the grid-stride layout from day 8, unrolled: a thread
// takes elements base, base + stride, base + 2 * stride, and so on, so every
// one of the kWide loads is a coalesced 128-byte span per warp.
//
// The three loops are separate on purpose. Loading all kWide values, then
// transforming all kWide, then storing all kWide, is what puts kWide
// requests in flight per thread, and it is also what makes the compiler hold
// kWide floats live across the middle loop. Both halves of that trade are
// the point: more work in flight per thread, fewer threads that fit.
//
// Launch assumption: gridDim.x * blockDim.x * kWide >= n. With that, the
// union of the indices below covers [0, n) exactly once.
// snippet: fused-wide
template <int kWide>
__global__ void chainFusedWide(const float* __restrict__ x,
const float* __restrict__ r,
float* __restrict__ out, size_t n) {
const size_t base =
blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
const size_t stride = gridDim.x * static_cast<size_t>(blockDim.x);
float v[kWide];
#pragma unroll
for (int k = 0; k < kWide; ++k) {
const size_t i = base + static_cast<size_t>(k) * stride;
v[k] = (i < n) ? scaleBias(x[i]) : 0.0f;
}
#pragma unroll
for (int k = 0; k < kWide; ++k) {
const size_t i = base + static_cast<size_t>(k) * stride;
v[k] = addResidual(gelu(v[k]), (i < n) ? r[i] : 0.0f);
}
#pragma unroll
for (int k = 0; k < kWide; ++k) {
const size_t i = base + static_cast<size_t>(k) * stride;
if (i < n) {
out[i] = clampRange(v[k]);
}
}
}
// end snippet: fused-wide
// The chain in double on the host. Written for obvious correctness, not
// speed: plain loop, no OpenMP, no intrinsics. It never allocates; the
// caller owns every buffer.
//
// The float constants promote to double rather than being respelled as
// double literals, so the reference checks the kernel's arithmetic and not
// the author's choice of how many digits to type.
static void chainCpu(const float* x, const float* r, double* out, size_t n) {
for (size_t i = 0; i < n; ++i) {
double v = static_cast<double>(x[i]) * kGamma + kBeta;
const double inner = kGeluA * (v + kGeluB * v * v * v);
v = 0.5 * v * (1.0 + std::tanh(inner));
v = v + static_cast<double>(r[i]);
if (v < kLo) {
v = kLo;
}
if (v > kHi) {
v = kHi;
}
out[i] = v;
}
}
// Returns the first index where got and want differ by more than the
// tolerance, or n if they agree everywhere. Returning the index rather than
// a bool is the whole point: "wrong at 512" names the block, "wrong" does
// not.
static size_t firstMismatch(const float* got, const double* want, size_t n) {
for (size_t i = 0; i < n; ++i) {
const double scale = std::fabs(want[i]);
const double allowed = kAbsTolerance + kRelTolerance * scale;
if (std::fabs(static_cast<double>(got[i]) - want[i]) > allowed) {
return i;
}
}
return n;
}
// Times a launch with CUDA events and returns the mean milliseconds per run.
// A host-side clock around a launch measures the launch, not the kernel,
// because launches are asynchronous. Day 9 takes that apart.
//
// This is the one template and the one lambda allowed in module 1 to 3 code.
// Copy it verbatim.
template <typename LaunchFn>
static float timeKernel(LaunchFn launch) {
cudaEvent_t start, stop;
CUDA_CHECK(cudaEventCreate(&start));
CUDA_CHECK(cudaEventCreate(&stop));
// Warm up this kernel, not just the first kernel in the program. Lazy
// module loading has been the default since CUDA 12.2 on Linux, so the
// first launch of each kernel pays its own load.
for (int i = 0; i < kWarmupRuns; ++i) {
launch();
}
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaGetLastError());
CUDA_CHECK(cudaEventRecord(start));
for (int i = 0; i < kTimedRuns; ++i) {
launch();
}
CUDA_CHECK(cudaEventRecord(stop));
CUDA_CHECK(cudaEventSynchronize(stop));
CUDA_CHECK(cudaGetLastError());
float ms = 0.0f;
CUDA_CHECK(cudaEventElapsedTime(&ms, start, stop));
CUDA_CHECK(cudaEventDestroy(start));
CUDA_CHECK(cudaEventDestroy(stop));
return ms / kTimedRuns;
}
static int blocksFor(size_t n) {
return static_cast<int>((n + kThreadsPerBlock - 1) / kThreadsPerBlock);
}
// Effective bandwidth from the compulsory byte count, not from a profiler.
// This is what the program can know: the bytes the algorithm has to move.
// What actually crossed the bus is dram__bytes in the shipped Nsight Compute
// report, and the gap between the two is the caches.
static double bandwidthGBs(double floatsPerElem, size_t n, float ms) {
const double bytes = floatsPerElem * static_cast<double>(n) * sizeof(float);
return bytes / (static_cast<double>(ms) * 1.0e-3) / 1.0e9;
}
// The four staged launches, in order, on the default stream. Enqueued and
// not synchronized: the caller decides where the measurement ends.
static void launchStaged(const float* d_x, const float* d_r, float* d_t1,
float* d_t2, float* d_t3, float* d_y, size_t n) {
const int blocks = blocksFor(n);
applyScaleBias<<<blocks, kThreadsPerBlock>>>(d_x, d_t1, n);
applyGelu<<<blocks, kThreadsPerBlock>>>(d_t1, d_t2, n);
applyAddResidual<<<blocks, kThreadsPerBlock>>>(d_t2, d_r, d_t3, n);
applyClampRange<<<blocks, kThreadsPerBlock>>>(d_t3, d_y, n);
}
struct WideRow {
int wide;
int blocks;
int regs;
size_t spillBytes;
int blocksPerSm;
double occupancy;
float ms;
size_t bad;
};
// Runs one width twice: once at the correctness size, where the result is
// compared against the double reference, and once at the timed size, where
// it is measured. Registers and theoretical occupancy come from the runtime
// API rather than from a profiler, so this table needs no performance
// counters and no root.
template <int kWide>
static void probeWide(WideRow* row, const float* d_x, const float* d_r,
float* d_out, float* h_scratch, const double* h_want,
int maxThreadsPerSm) {
const size_t perBlock =
static_cast<size_t>(kWide) * static_cast<size_t>(kThreadsPerBlock);
const int checkBlocks =
static_cast<int>((kCheckElems + perBlock - 1) / perBlock);
const int blocks = static_cast<int>((kElems + perBlock - 1) / perBlock);
chainFusedWide<kWide>
<<<checkBlocks, kThreadsPerBlock>>>(d_x, d_r, d_out, kCheckElems);
CUDA_CHECK(cudaGetLastError());
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaMemcpy(h_scratch, d_out, kCheckElems * sizeof(float),
cudaMemcpyDeviceToHost));
row->bad = firstMismatch(h_scratch, h_want, kCheckElems);
cudaFuncAttributes attr;
CUDA_CHECK(cudaFuncGetAttributes(&attr, chainFusedWide<kWide>));
int blocksPerSm = 0;
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&blocksPerSm, chainFusedWide<kWide>, kThreadsPerBlock, 0));
row->wide = kWide;
row->blocks = blocks;
row->regs = attr.numRegs;
row->spillBytes = attr.localSizeBytes;
row->blocksPerSm = blocksPerSm;
row->occupancy = 100.0 * static_cast<double>(blocksPerSm) *
kThreadsPerBlock / static_cast<double>(maxThreadsPerSm);
row->ms = timeKernel([&] {
chainFusedWide<kWide>
<<<blocks, kThreadsPerBlock>>>(d_x, d_r, d_out, kElems);
});
}
// Runs one kernel or chain at the correctness size and compares the result
// against the double reference. Returns the first bad index, or kCheckElems
// when everything matched.
static size_t checkAtCheckSize(bool fused, const float* d_x, const float* d_r,
float* d_t1, float* d_t2, float* d_t3,
float* d_y, float* h_scratch,
const double* h_want) {
if (fused) {
chainFused<<<blocksFor(kCheckElems), kThreadsPerBlock>>>(d_x, d_r, d_y,
kCheckElems);
} else {
launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, kCheckElems);
}
CUDA_CHECK(cudaGetLastError());
CUDA_CHECK(cudaDeviceSynchronize());
CUDA_CHECK(cudaMemcpy(h_scratch, d_y, kCheckElems * sizeof(float),
cudaMemcpyDeviceToHost));
return firstMismatch(h_scratch, h_want, kCheckElems);
}
int main() {
const int device = 0;
CUDA_CHECK(cudaSetDevice(device));
cudaDeviceProp prop;
CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
const size_t bytes = kElems * sizeof(float);
const double mib = static_cast<double>(bytes) / (1024.0 * 1024.0);
// Registers per thread at which this card still fits its full complement
// of resident threads on one SM. Derived from the card, not asserted:
// the number is 64 on a T4 and it is not 64 everywhere.
const int regsAtFullOccupancy =
prop.regsPerMultiprocessor / prop.maxThreadsPerMultiProcessor;
std::vector<float> h_x(kElems);
std::vector<float> h_r(kElems);
std::vector<float> h_staged(kElems);
std::vector<float> h_fused(kElems);
std::vector<double> h_want(kCheckElems);
// x sweeps [-4, 4) so that stage 1 produces [-6.25, 5.75) and the clamp
// at kHi = 3 catches a real share of the elements rather than none.
for (size_t i = 0; i < kElems; ++i) {
h_x[i] = (static_cast<float>(i % 2048) / 1024.0f - 1.0f) * 4.0f;
h_r[i] = static_cast<float>(i % 97) / 97.0f - 0.5f;
}
chainCpu(h_x.data(), h_r.data(), h_want.data(), kCheckElems);
// How much of the reference sits on a clamp. A stage that never fires is
// a stage that might not be there.
size_t clamped = 0;
for (size_t i = 0; i < kCheckElems; ++i) {
if (h_want[i] == static_cast<double>(kHi) ||
h_want[i] == static_cast<double>(kLo)) {
++clamped;
}
}
float* d_x = nullptr;
float* d_r = nullptr;
float* d_t1 = nullptr;
float* d_t2 = nullptr;
float* d_t3 = nullptr;
float* d_y = nullptr;
CUDA_CHECK(cudaMalloc(&d_x, bytes));
CUDA_CHECK(cudaMalloc(&d_r, bytes));
CUDA_CHECK(cudaMalloc(&d_t1, bytes));
CUDA_CHECK(cudaMalloc(&d_t2, bytes));
CUDA_CHECK(cudaMalloc(&d_t3, bytes));
CUDA_CHECK(cudaMalloc(&d_y, bytes));
CUDA_CHECK(cudaMemcpy(d_x, h_x.data(), bytes, cudaMemcpyHostToDevice));
CUDA_CHECK(cudaMemcpy(d_r, h_r.data(), bytes, cudaMemcpyHostToDevice));
int status = EXIT_SUCCESS;
const size_t badStaged = checkAtCheckSize(
false, d_x, d_r, d_t1, d_t2, d_t3, d_y, h_staged.data(), h_want.data());
const size_t badFused = checkAtCheckSize(
true, d_x, d_r, d_t1, d_t2, d_t3, d_y, h_fused.data(), h_want.data());
// Real branches returning EXIT_FAILURE, not assert(). CI builds Release,
// Release defines NDEBUG, and NDEBUG deletes assert(), so a check written
// that way would vanish in exactly the build that matters.
if (badStaged != kCheckElems) {
std::fprintf(stderr, "staged chain wrong at %zu: got %.9g, want %.9g\n",
badStaged, static_cast<double>(h_staged[badStaged]),
h_want[badStaged]);
status = EXIT_FAILURE;
}
if (badFused != kCheckElems) {
std::fprintf(stderr, "chainFused wrong at %zu: got %.9g, want %.9g\n",
badFused, static_cast<double>(h_fused[badFused]),
h_want[badFused]);
status = EXIT_FAILURE;
}
float copyMs = 0.0f;
float stageMs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
float stagedChainMs = 0.0f;
float fusedChainMs = 0.0f;
double launchUsSmall = 0.0;
double launchUsFull = 0.0;
float sweepStagedMs[kSweepCount] = {0.0f};
float sweepFusedMs[kSweepCount] = {0.0f};
WideRow wide[kWideCount] = {};
size_t differing = 0;
double maxAbsDiff = 0.0;
if (status == EXIT_SUCCESS) {
// Part 1. A batch of empty launches, back to back on one stream.
// With the host free to run ahead, what this measures is the gap the
// device leaves between one kernel ending and the next beginning,
// which is the cost a chain of dependent kernels actually pays. A
// single cold launch costs more; day 9 measured that one.
const float smallBatchMs = timeKernel([&] {
for (int k = 0; k < kLaunchesPerBatch; ++k) {
doNothing<<<1, 1>>>();
}
});
const float fullBatchMs = timeKernel([&] {
for (int k = 0; k < kLaunchesPerBatch; ++k) {
doNothing<<<blocksFor(kElems), kThreadsPerBlock>>>();
}
});
launchUsSmall =
static_cast<double>(smallBatchMs) * 1000.0 / kLaunchesPerBatch;
launchUsFull =
static_cast<double>(fullBatchMs) * 1000.0 / kLaunchesPerBatch;
// Part 2. Every timed launch reads one buffer and writes another and
// the two are never swapped, so each kernel sees the same input on
// every run and no reset has to sit inside the measurement.
const int blocks = blocksFor(kElems);
copyMs = timeKernel(
[&] { copyFloor<<<blocks, kThreadsPerBlock>>>(d_x, d_y, kElems); });
stageMs[0] = timeKernel([&] {
applyScaleBias<<<blocks, kThreadsPerBlock>>>(d_x, d_t1, kElems);
});
stageMs[1] = timeKernel([&] {
applyGelu<<<blocks, kThreadsPerBlock>>>(d_t1, d_t2, kElems);
});
stageMs[2] = timeKernel([&] {
applyAddResidual<<<blocks, kThreadsPerBlock>>>(d_t2, d_r, d_t3,
kElems);
});
stageMs[3] = timeKernel([&] {
applyClampRange<<<blocks, kThreadsPerBlock>>>(d_t3, d_y, kElems);
});
stagedChainMs = timeKernel(
[&] { launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, kElems); });
CUDA_CHECK(
cudaMemcpy(h_staged.data(), d_y, bytes, cudaMemcpyDeviceToHost));
fusedChainMs = timeKernel([&] {
chainFused<<<blocks, kThreadsPerBlock>>>(d_x, d_r, d_y, kElems);
});
CUDA_CHECK(
cudaMemcpy(h_fused.data(), d_y, bytes, cudaMemcpyDeviceToHost));
// Does fusion change the answer? Both versions round to float at the
// same four points, so they should agree exactly, but the compiler
// is free to contract a multiply and an add across a stage boundary
// in the fused kernel and not in the staged one. Counted rather than
// assumed, and reported rather than gated: a difference here is a
// fact about fusion, not a failure.
for (size_t i = 0; i < kElems; ++i) {
if (h_staged[i] != h_fused[i]) {
++differing;
const double d = std::fabs(static_cast<double>(h_staged[i]) -
static_cast<double>(h_fused[i]));
if (d > maxAbsDiff) {
maxAbsDiff = d;
}
}
}
// Part 3. The same two chains, smaller and smaller, until four
// launches cost more than the bytes they save.
for (int s = 0; s < kSweepCount; ++s) {
const size_t n = kSweep[s];
const int sweepBlocks = blocksFor(n);
sweepStagedMs[s] = timeKernel(
[&] { launchStaged(d_x, d_r, d_t1, d_t2, d_t3, d_y, n); });
sweepFusedMs[s] = timeKernel([&] {
chainFused<<<sweepBlocks, kThreadsPerBlock>>>(d_x, d_r, d_y, n);
});
}
// Part 4. Six widths, instantiated by hand because a template
// argument cannot come from a loop. The static_assert above keeps
// this list and kWideK in step.
//
// h_staged is reused as the correctness scratch buffer here. The
// staged-against-fused comparison above has already run, so nothing
// still needs what it held.
probeWide<1>(&wide[0], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
probeWide<2>(&wide[1], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
probeWide<4>(&wide[2], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
probeWide<8>(&wide[3], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
probeWide<16>(&wide[4], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
probeWide<32>(&wide[5], d_x, d_r, d_y, h_staged.data(), h_want.data(),
prop.maxThreadsPerMultiProcessor);
for (int w = 0; w < kWideCount; ++w) {
if (wide[w].bad != kCheckElems) {
std::fprintf(stderr,
"chainFusedWide<%d> wrong at %zu, want %.9g\n",
wide[w].wide, wide[w].bad, h_want[wide[w].bad]);
status = EXIT_FAILURE;
}
}
}
// Arithmetic, not a claim about the card: a zero or negative elapsed time
// means the event pair never separated, and every figure derived from it
// would be infinite or nonsense.
if (status == EXIT_SUCCESS) {
if (copyMs <= 0.0f || stagedChainMs <= 0.0f || fusedChainMs <= 0.0f ||
launchUsSmall <= 0.0 || launchUsFull <= 0.0) {
std::fprintf(stderr,
"timing broken: copy %.6f, staged %.6f, fused %.6f, "
"launch %.6f us\n",
static_cast<double>(copyMs),
static_cast<double>(stagedChainMs),
static_cast<double>(fusedChainMs), launchUsSmall);
status = EXIT_FAILURE;
}
for (int s = 0; s < kSweepCount; ++s) {
if (sweepStagedMs[s] <= 0.0f || sweepFusedMs[s] <= 0.0f) {
std::fprintf(stderr, "timing broken in the sweep at n = %zu\n",
kSweep[s]);
status = EXIT_FAILURE;
}
}
for (int w = 0; w < kWideCount; ++w) {
if (wide[w].ms <= 0.0f) {
std::fprintf(stderr, "timing broken at width %d\n",
wide[w].wide);
status = EXIT_FAILURE;
}
}
}
// The staged chain has to be slower than the fused one on this hardware
// and at this size, because it does the same arithmetic and moves three
// times the bytes. If it is not, the two are not running the same chain
// and no ratio on the page means anything.
if (status == EXIT_SUCCESS && stagedChainMs <= fusedChainMs) {
std::fprintf(stderr,
"the staged chain (%.6f ms) is not slower than the fused "
"one (%.6f ms), which cannot be right at 3x the traffic\n",
static_cast<double>(stagedChainMs),
static_cast<double>(fusedChainMs));
status = EXIT_FAILURE;
}
if (status == EXIT_SUCCESS) {
std::printf("GPU: %s (compute capability %d.%d, %d SMs)\n", prop.name,
prop.major, prop.minor, prop.multiProcessorCount);
std::printf(
"%d resident threads/SM, %d registers/SM, %d KiB L2, %d "
"threads/block\n",
prop.maxThreadsPerMultiProcessor, prop.regsPerMultiprocessor,
prop.l2CacheSize / 1024, kThreadsPerBlock);
std::printf(
"Timed size %zu elements, %.1f MiB per buffer, six buffers\n",
kElems, mib);
std::printf("Correctness size %zu elements, %.1f%% of them clamped\n",
kCheckElems,
100.0 * static_cast<double>(clamped) /
static_cast<double>(kCheckElems));
std::printf(
" staged chain, chainFused and all %d widths match the CPU\n"
" reference at that size, which no block size divides\n\n",
kWideCount);
std::printf("Bytes per element, from the constants\n");
std::printf(" staged chain %5.0f (2 + 2 + 3 + 2 floats)\n",
kStagedFloatsPerElem * sizeof(float));
std::printf(" chainFused %5.0f (x, residual, out)\n",
kFusedFloatsPerElem * sizeof(float));
std::printf(" ratio %5.2f\n\n",
kStagedFloatsPerElem / kFusedFloatsPerElem);
std::printf(
"Part 1: what one launch costs, %d back-to-back launches per "
"timed run\n",
kLaunchesPerBatch);
std::printf("%-38s %16s\n", "launch shape", "per launch (us)");
std::printf("%-38s %16s\n", "-------------------------------------",
"---------------");
std::printf("%-38s %16.3f\n", "doNothing, 1 block of 1 thread",
launchUsSmall);
std::printf("%-38s %16.3f\n\n", "doNothing, the chain's own grid",
launchUsFull);
std::printf(
"Part 2: the chain at %zu elements, mean of %d runs after %d "
"warm-ups\n",
kElems, kTimedRuns, kWarmupRuns);
std::printf("%-28s %8s %11s %9s %8s\n", "kernel or chain", "ms",
"bytes/elem", "GB/s", "x copy");
std::printf("%-28s %8s %11s %9s %8s\n", "---------------------------",
"-------", "----------", "--------", "-------");
// Every row's effective bandwidth is its own compulsory bytes over
// its own time, and the last column divides that by the floor row's,
// so two rows moving different byte counts are still comparable.
const double copyGBs = bandwidthGBs(kCopyFloatsPerElem, kElems, copyMs);
std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", "copyFloor (the floor)",
static_cast<double>(copyMs),
kCopyFloatsPerElem * sizeof(float), copyGBs, 1.00);
const char* stageNames[4] = {"applyScaleBias", "applyGelu",
"applyAddResidual", "applyClampRange"};
const double stageFloats[4] = {2.0, 2.0, 3.0, 2.0};
double summed = 0.0;
for (int s = 0; s < 4; ++s) {
summed += static_cast<double>(stageMs[s]);
const double gbs = bandwidthGBs(stageFloats[s], kElems, stageMs[s]);
std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", stageNames[s],
static_cast<double>(stageMs[s]),
stageFloats[s] * sizeof(float), gbs, gbs / copyGBs);
}
const double summedGBs = bandwidthGBs(kStagedFloatsPerElem, kElems,
static_cast<float>(summed));
std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n", " the four, summed",
summed, kStagedFloatsPerElem * sizeof(float), summedGBs,
summedGBs / copyGBs);
const double stagedGBs =
bandwidthGBs(kStagedFloatsPerElem, kElems, stagedChainMs);
std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n",
"staged chain, timed as one",
static_cast<double>(stagedChainMs),
kStagedFloatsPerElem * sizeof(float), stagedGBs,
stagedGBs / copyGBs);
const double fusedGBs =
bandwidthGBs(kFusedFloatsPerElem, kElems, fusedChainMs);
std::printf("%-28s %8.3f %11.0f %9.1f %8.2f\n\n", "chainFused",
static_cast<double>(fusedChainMs),
kFusedFloatsPerElem * sizeof(float), fusedGBs,
fusedGBs / copyGBs);
std::printf("staged / fused = %.2f, and the byte ratio is %.2f\n",
static_cast<double>(stagedChainMs / fusedChainMs),
kStagedFloatsPerElem / kFusedFloatsPerElem);
if (differing == 0) {
std::printf(
"The two chains agree bit for bit on all %zu elements.\n\n",
kElems);
} else {
std::printf(
"The two chains differ on %zu of %zu elements, at most by "
"%.3g.\n\n",
differing, kElems, maxAbsDiff);
}
std::printf("Part 3: the same two chains, swept over n\n");
std::printf("%11s %12s %12s %13s %17s\n", "n", "staged (ms)",
"fused (ms)", "staged/fused", "staged us/launch");
std::printf("%11s %12s %12s %13s %17s\n", "----------", "-----------",
"-----------", "------------", "----------------");
for (int s = 0; s < kSweepCount; ++s) {
std::printf("%11zu %12.4f %12.4f %13.2f %17.3f\n", kSweep[s],
static_cast<double>(sweepStagedMs[s]),
static_cast<double>(sweepFusedMs[s]),
static_cast<double>(sweepStagedMs[s] / sweepFusedMs[s]),
static_cast<double>(sweepStagedMs[s]) * 1000.0 / 4.0);
}
std::printf(
"\nThe last column is the staged chain's time divided by its four\n"
"launches. As n falls it has to approach part 1's per-launch "
"figure,\nbecause that is all the chain is still doing.\n\n");
std::printf(
"Part 4: chainFused widened, same work and same bytes every "
"row\n");
std::printf("%11s %8s %6s %8s %10s %10s %8s %9s\n", "per thread",
"blocks", "regs", "spill B", "blocks/SM", "occupancy", "ms",
"GB/s");
std::printf("%11s %8s %6s %8s %10s %10s %8s %9s\n", "----------",
"-------", "-----", "-------", "---------", "---------",
"-------", "--------");
for (int w = 0; w < kWideCount; ++w) {
std::printf("%11d %8d %6d %8zu %10d %9.1f%% %8.3f %9.1f\n",
wide[w].wide, wide[w].blocks, wide[w].regs,
wide[w].spillBytes, wide[w].blocksPerSm,
wide[w].occupancy, static_cast<double>(wide[w].ms),
bandwidthGBs(kFusedFloatsPerElem, kElems, wide[w].ms));
}
std::printf(
"\nThis card keeps all %d of its resident threads on an SM only "
"while a\nthread uses %d registers or fewer, which is %d "
"registers per SM\ndivided by %d threads. Above that, warps come "
"off the SM.\n",
prop.maxThreadsPerMultiProcessor, regsAtFullOccupancy,
prop.regsPerMultiprocessor, prop.maxThreadsPerMultiProcessor);
}
CUDA_CHECK(cudaFree(d_x));
CUDA_CHECK(cudaFree(d_r));
CUDA_CHECK(cudaFree(d_t1));
CUDA_CHECK(cudaFree(d_t2));
CUDA_CHECK(cudaFree(d_t3));
CUDA_CHECK(cudaFree(d_y));
return status;
}