COURSE / SOURCE

mma_sync.cu

All lessons
Source filecode/day73-mma-sync/mma_sync.cu

This 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 73: mma.sync and ldmatrix. One m16n8k8 tensor-core tile written by
// hand in inline PTX, first alone so every lane's fragment ownership is
// visible, then inside a full 2048x2048 FP16-in FP32-accumulate matmul.
//
// The fragment formulas below are transcribed from the PTX ISA 8.5 tables
// ("Matrix Fragments for mma.m16n8k8", CUDA 12.6 docs) and the program
// checks them twice. The probe gives every cell of its tiles a unique
// value and decodes the loaded registers back into coordinates, which
// tests the two ldmatrix calls on their own. The matmul then gates the
// product exactly, which tests the accumulator formula, whose row and
// column rule is identical to A's and so cannot be told apart by an
// ownership map alone.
//
// Matmul inputs are multiples of 0.25 in [-2, 2], so every product is a
// multiple of 1/16 and every FP32 partial sum is exact (|sum| * 16 <
// 2^24). The double-precision reference must match bit for bit.
//
// Build (T4, Colab): nvcc -std=c++17 -O3 -arch=sm_75 -o mma_sync mma_sync.cu
// Build (Ampere+):   nvcc -std=c++17 -O3 -arch=sm_80 -o mma_sync mma_sync.cu
// The m16n8k16 probe sits behind #if __CUDA_ARCH__ >= 800: the sm_75 build
// compiles it to an empty kernel and the host never launches it below 8.0.

#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <vector>

#include <cuda_fp16.h>
#include <cuda_runtime.h>

#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)

// The matmul: C (f32) = A (f16) x B (f16), all square, row major.
constexpr int kSize = 2048;
// Block tile: 64x64 output from 32-deep slices of A and B.
constexpr int kBlockM = 64;
constexpr int kBlockN = 64;
constexpr int kBlockK = 32;
// Tile rows are padded by 8 halves (16 bytes). The pad shifts consecutive
// rows to different banks; 16 bytes keeps every row start aligned for the
// uint4 fill and for ldmatrix, which loads 16-byte row segments.
constexpr int kPad = 8;
constexpr int kLdA = kBlockK + kPad;  // 40 halves per A tile row
constexpr int kLdB = kBlockN + kPad;  // 72 halves per B tile row
// Eight warps in a 4x2 grid; each owns a 16x32 slice of the output, which
// is one m16 row band times four n8 tiles.
constexpr int kWarpsM = 4;
constexpr int kWarpsN = 2;
constexpr int kWarpN = 32;
// The tile shape forces the block size: 8 warps of 32 (it is still 256).
constexpr int kThreadsPerBlock = kWarpsM * kWarpsN * 32;
constexpr int kWarmupRuns = 3;
constexpr int kTimedRuns = 10;

static_assert(kSize % kBlockM == 0 && kSize % kBlockN == 0 &&
                  kSize % kBlockK == 0,
              "the kernel has no edge guards; tiles must divide the size");
static_assert(kBlockK % 8 == 0 && kWarpsM * 16 == kBlockM &&
                  kWarpsN * kWarpN == kBlockN,
              "warp tiles must cover the block tile exactly");
static_assert(kLdA * 2 % 16 == 0 && kLdB * 2 % 16 == 0,
              "padded rows must stay 16-byte aligned for uint4 and ldmatrix");
static_assert(kBlockM * kBlockK / 8 == kThreadsPerBlock &&
                  kBlockK * kBlockN / 8 == kThreadsPerBlock,
              "the fill gives each thread exactly one uint4 per tile");

// Converts a shared-memory pointer to the 32-bit shared-space address the
// ldmatrix operand wants.
static __device__ unsigned smemAddr(const void* p) {
    return static_cast<unsigned>(__cvta_generic_to_shared(p));
}

// One element of an .f16x2 register, as a float. The PTX fragment tables
// number elements low to high, and this is what "low" means: the
// lower-numbered element sits in the low 16 bits.
static __device__ float halfOf(unsigned pair, int which) {
    const unsigned short bits = static_cast<unsigned short>(
        which != 0 ? (pair >> 16) : (pair & 0xffffu));
    return __half2float(__ushort_as_half(bits));
}

// Two 8x8 b16 matrices into two registers: the A fragment of one m16n8k8
// tile. Lanes 0-7 name rows 0-7, lanes 8-15 rows 8-15.
static __device__ void ldmatrixX2(unsigned& r0, unsigned& r1, unsigned addr) {
    asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0, %1}, [%2];"
                 : "=r"(r0), "=r"(r1)
                 : "r"(addr));
}

// One row-major 8x8 loaded column major: the B fragment. Lanes 0-7 name
// the eight rows; on sm_75 lanes 8-31 must still hold valid addresses.
static __device__ void ldmatrixX1Trans(unsigned& r0, unsigned addr) {
    asm volatile("ldmatrix.sync.aligned.m8n8.x1.trans.shared.b16 {%0}, [%1];"
                 : "=r"(r0)
                 : "r"(addr));
}

// snippet: ldmatrix-b
// Four row-major 8x8 tiles loaded column major in one instruction: the
// four B fragments a 16x32 warp tile needs per k step.
static __device__ void ldmatrixX4Trans(unsigned& r0, unsigned& r1, unsigned& r2,
                                       unsigned& r3, unsigned addr) {
    asm volatile(
        "ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 "
        "{%0, %1, %2, %3}, [%4];"
        : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
        : "r"(addr));
}
// end snippet

// snippet: mma-asm
// One tensor-core instruction: D (16x8, f32) += A (16x8, f16) x B (8x8,
// f16). The .row.col layout is the only one the shape offers, which is
// why the B loads above carry .trans.
static __device__ void mmaM16N8K8(float acc[4], unsigned a0, unsigned a1,
                                  unsigned b0) {
    asm volatile(
        "mma.sync.aligned.m16n8k8.row.col.f32.f16.f16.f32 "
        "{%0, %1, %2, %3}, {%4, %5}, {%6}, {%0, %1, %2, %3};"
        : "+f"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3])
        : "r"(a0), "r"(a1), "r"(b0));
}
// end snippet: mma-asm

// One warp, one m16n8k8 tile: d = a x b, and a copy of what ldmatrix put
// in each lane's registers.
//
// One thread holds two A registers, one B register and four D floats. The
// four halves of the A pair and the two halves of the B register are
// written out as floats, so the host can name the (row, col) each lane
// received rather than take the PTX table's word for it. Storing the
// store indices instead would prove nothing: they come from the same
// formula the store uses.
// One warp's addresses: the scalar fills are contiguous; each ldmatrix
// reads sixteen-byte row segments, one per four-lane group.
// Launch: exactly one block of 32 threads.
__global__ void probeFragmentMap(const __half* __restrict__ a,
                                 const __half* __restrict__ b,
                                 float* __restrict__ d,
                                 float* __restrict__ aFrag,
                                 float* __restrict__ bFrag) {
    __shared__ __half tileA[16 * 8];
    __shared__ __half tileB[8 * 8];

    const unsigned int lane = threadIdx.x % 32u;

    for (int e = static_cast<int>(lane); e < 16 * 8; e += 32) {
        tileA[e] = a[e];
    }
    for (int e = static_cast<int>(lane); e < 8 * 8; e += 32) {
        tileB[e] = b[e];
    }
    __syncwarp();

    // On sm_75 every lane must hand ldmatrix a valid address, including
    // the lanes above the ones whose rows are used ("For .target sm_75 or
    // below, all threads must contain valid addresses", PTX ISA 8.5,
    // ldmatrix). The modulus below is that rule, not a convenience.
    unsigned a0, a1;
    ldmatrixX2(a0, a1, smemAddr(&tileA[(lane % 16u) * 8]));
    unsigned b0;
    ldmatrixX1Trans(b0, smemAddr(&tileB[(lane % 8u) * 8]));

    float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
    mmaM16N8K8(acc, a0, a1, b0);

    // snippet: fragment-report
    // a0 carries elements a0 and a1 of the fragment, a1 carries a2 and a3.
    aFrag[lane * 4 + 0] = halfOf(a0, 0);
    aFrag[lane * 4 + 1] = halfOf(a0, 1);
    aFrag[lane * 4 + 2] = halfOf(a1, 0);
    aFrag[lane * 4 + 3] = halfOf(a1, 1);
    bFrag[lane * 2 + 0] = halfOf(b0, 0);
    bFrag[lane * 2 + 1] = halfOf(b0, 1);
    // end snippet: fragment-report

    // snippet: accumulator-store
    // The accumulator table for m16n8k8: PTX calls these groupID and
    // threadID_in_group. c0 and c1 sit in row groupID, c2 and c3 eight
    // rows below, and each four-lane group covers one row pair.
    const unsigned int group = lane / 4u;
    const unsigned int member = lane % 4u;
    for (int i = 0; i < 4; ++i) {
        const unsigned int row = group + ((i < 2) ? 0u : 8u);
        const unsigned int col = member * 2u + (i & 1);
        d[row * 8 + col] = acc[i];
    }
    // end snippet: accumulator-store
}

// c = a x b at kSize, FP16 in, FP32 accumulate, one m16n8k8 tile at a
// time.
//
// One thread carries 16 accumulator floats: its share of the four n8
// tiles its warp owns. One warp's ldmatrix addresses name 16-byte row
// segments of the padded shared tiles; the uint4 fill before them is
// contiguous, eight halves per thread per tile.
// Launch: kThreadsPerBlock threads, grid (kSize / kBlockN, kSize /
// kBlockM); the static_asserts pin every divisibility this relies on.
__global__ void matmulMmaSync(const __half* __restrict__ a,
                              const __half* __restrict__ b,
                              float* __restrict__ c, int n) {
    __shared__ __half tileA[kBlockM * kLdA];
    __shared__ __half tileB[kBlockK * kLdB];

    const unsigned int tid = threadIdx.x;
    const unsigned int warp = tid / 32u;
    const unsigned int lane = tid % 32u;
    const unsigned int warpM = warp % kWarpsM;
    const unsigned int warpN = warp / kWarpsM;

    const int blockRow = static_cast<int>(blockIdx.y) * kBlockM;
    const int blockCol = static_cast<int>(blockIdx.x) * kBlockN;

    float acc[4][4] = {};

    // Fragment addresses, fixed for the whole k loop. A: lanes 0-15 name
    // the sixteen rows of the warp's 16x8 slice; lanes 16-31 repeat them
    // because sm_75 requires a valid address in every lane. B: each group
    // of eight lanes names the eight rows of one 8x8 tile, tiles stepping
    // right by eight columns.
    const unsigned int rowA = warpM * 16u + (lane % 16u);
    const unsigned int rowB = lane % 8u;
    const unsigned int colB = warpN * kWarpN + (lane / 8u) * 8u;

    for (int kt = 0; kt < n; kt += kBlockK) {
        const int e = static_cast<int>(tid);
        const int rA = e / (kBlockK / 8);
        const int cA = (e % (kBlockK / 8)) * 8;
        *reinterpret_cast<uint4*>(&tileA[rA * kLdA + cA]) =
            *reinterpret_cast<const uint4*>(
                &a[static_cast<size_t>(blockRow + rA) * n + kt + cA]);
        const int rB = e / (kBlockN / 8);
        const int cB = (e % (kBlockN / 8)) * 8;
        *reinterpret_cast<uint4*>(&tileB[rB * kLdB + cB]) =
            *reinterpret_cast<const uint4*>(
                &b[static_cast<size_t>(kt + rB) * n + blockCol + cB]);
        __syncthreads();

        // snippet: inner-loop
        for (int ks = 0; ks < kBlockK; ks += 8) {
            unsigned a0, a1;
            ldmatrixX2(a0, a1, smemAddr(&tileA[rowA * kLdA + ks]));
            unsigned b0, b1, b2, b3;
            ldmatrixX4Trans(b0, b1, b2, b3,
                            smemAddr(&tileB[(ks + rowB) * kLdB + colB]));
            mmaM16N8K8(acc[0], a0, a1, b0);
            mmaM16N8K8(acc[1], a0, a1, b1);
            mmaM16N8K8(acc[2], a0, a1, b2);
            mmaM16N8K8(acc[3], a0, a1, b3);
        }
        // end snippet
        __syncthreads();
    }

    // The same accumulator formulas probeFragmentMap prints.
    const unsigned int group = lane / 4u;
    const unsigned int member = lane % 4u;
    for (int j = 0; j < 4; ++j) {
        for (int i = 0; i < 4; ++i) {
            const int row = blockRow + static_cast<int>(warpM) * 16 +
                            static_cast<int>(group) + ((i < 2) ? 0 : 8);
            const int col = blockCol + static_cast<int>(warpN) * kWarpN +
                            j * 8 + static_cast<int>(member) * 2 + (i & 1);
            c[static_cast<size_t>(row) * n + col] = acc[j][i];
        }
    }
}

// One warp, one m16n8k16 tile: the sm_80 shape, twice the K depth per
// instruction. The sm_75 build compiles this to an empty kernel and the
// host never launches it below compute capability 8.0.
//
// One thread holds four A registers (eight halves) and two B registers.
// One warp's ldmatrix.x4 addresses walk the four 8x8 quadrants of A in
// the fragment table's order: rows 0-7 then 8-15 of columns 0-7, then the
// same two bands of columns 8-15.
// Launch: exactly one block of 32 threads.
__global__ void probeTileK16(const __half* __restrict__ a,
                             const __half* __restrict__ b,
                             float* __restrict__ d) {
#if __CUDA_ARCH__ >= 800
    __shared__ __half tileA[16 * 16];
    __shared__ __half tileB[16 * 8];

    const unsigned int lane = threadIdx.x % 32u;

    for (int e = static_cast<int>(lane); e < 16 * 16; e += 32) {
        tileA[e] = a[e];
    }
    for (int e = static_cast<int>(lane); e < 16 * 8; e += 32) {
        tileB[e] = b[e];
    }
    __syncwarp();

    const unsigned int quad = lane / 8u;
    const unsigned int rowQ = lane % 8u;
    const unsigned int rA = rowQ + ((quad == 1u || quad == 3u) ? 8u : 0u);
    const unsigned int cA = (quad >= 2u) ? 8u : 0u;
    unsigned a0, a1, a2, a3;
    asm volatile(
        "ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];"
        : "=r"(a0), "=r"(a1), "=r"(a2), "=r"(a3)
        : "r"(smemAddr(&tileA[rA * 16 + cA])));

    // B is 16x8: two stacked 8x8 tiles, loaded column major.
    const unsigned int rB = (lane % 8u) + ((lane / 8u) % 2u) * 8u;
    unsigned b0, b1;
    asm volatile(
        "ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0, %1}, [%2];"
        : "=r"(b0), "=r"(b1)
        : "r"(smemAddr(&tileB[rB * 8])));

    float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f};
    // snippet: k16-mma
    asm volatile(
        "mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
        "{%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {%0, %1, %2, %3};"
        : "+f"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3])
        : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
    // end snippet

    const unsigned int group = lane / 4u;
    const unsigned int member = lane % 4u;
    for (int i = 0; i < 4; ++i) {
        const unsigned int row = group + ((i < 2) ? 0u : 8u);
        const unsigned int col = member * 2u + (i & 1);
        d[row * 8 + col] = acc[i];
    }
#endif
}

// CPU reference. Plain loops, double accumulate, no allocation. With this
// program's exact inputs the cast at the end loses nothing.
static void matmulCpu(const float* a, const float* b, float* out, int m, int n,
                      int k) {
    for (int row = 0; row < m; ++row) {
        for (int col = 0; col < n; ++col) {
            double sum = 0.0;
            for (int p = 0; p < k; ++p) {
                sum += static_cast<double>(a[row * k + p]) *
                       static_cast<double>(b[p * n + col]);
            }
            out[static_cast<size_t>(row) * n + col] = static_cast<float>(sum);
        }
    }
}

// Returns the first index where got and want differ at all, or n if they
// agree everywhere. The gate is exact, not a tolerance: in both phases
// every product and every partial sum is exactly representable in FP32,
// so the only thing a nonzero difference can mean is a wrong index.
static size_t firstMismatch(const float* got, const float* want, size_t n) {
    for (size_t i = 0; i < n; ++i) {
        if (got[i] != want[i]) {
            return i;
        }
    }
    return n;
}

// Turns a probe value back into the (row, col) of the cell that holds it.
// Out-of-range values report (-1, -1), which is what a scrambled load
// looks like on the page rather than a silent wrong coordinate.
static void decodeCell(float value, float base, int cols, int cells, int* row,
                       int* col) {
    const float offset = value - base;
    const int index = static_cast<int>(offset);
    if (offset != static_cast<float>(index) || index < 0 || index >= cells) {
        *row = -1;
        *col = -1;
        return;
    }
    *row = index / cols;
    *col = index % cols;
}

// Times a launch with CUDA events and returns the mean milliseconds per run.
//
// This is the one template and the one lambda allowed in module 1 to 3 code.
// Copy it verbatim; the alternative is six copies of the event boilerplate,
// which is how a warm-up goes missing from one of them.
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;
}

// Exact-in-f16 test values: multiples of 0.25 in [-2, 2]. The two moduli
// are coprime to the tile sizes so no tile repeats another.
static float valueA(size_t i) {
    return static_cast<float>(static_cast<int>(i % 17) - 8) * 0.25f;
}
static float valueB(size_t i) {
    return static_cast<float>(static_cast<int>(i * 7 % 13) - 6) * 0.25f;
}

int main() {
    const int device = 0;
    CUDA_CHECK(cudaSetDevice(device));
    cudaDeviceProp prop;
    CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
    std::printf("GPU: %s (compute capability %d.%d)\n", prop.name, prop.major,
                prop.minor);

    const size_t elems = static_cast<size_t>(kSize) * kSize;
    std::vector<__half> h_a(elems);
    std::vector<__half> h_b(elems);
    std::vector<float> h_af(elems);
    std::vector<float> h_bf(elems);
    for (size_t i = 0; i < elems; ++i) {
        // __float2half is a documented host-side conversion helper; the
        // round trip through h_af keeps the reference fed with exactly
        // the values the GPU sees.
        h_a[i] = __float2half(valueA(i));
        h_b[i] = __float2half(valueB(i));
        h_af[i] = __half2float(h_a[i]);
        h_bf[i] = __half2float(h_b[i]);
    }

    __half* d_a = nullptr;
    __half* d_b = nullptr;
    float* d_c = nullptr;
    float* d_aFrag = nullptr;
    float* d_bFrag = nullptr;
    CUDA_CHECK(cudaMalloc(&d_a, elems * sizeof(__half)));
    CUDA_CHECK(cudaMalloc(&d_b, elems * sizeof(__half)));
    CUDA_CHECK(cudaMalloc(&d_c, elems * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_aFrag, 32 * 4 * sizeof(float)));
    CUDA_CHECK(cudaMalloc(&d_bFrag, 32 * 2 * sizeof(float)));
    CUDA_CHECK(cudaMemcpy(d_a, h_a.data(), elems * sizeof(__half),
                          cudaMemcpyHostToDevice));
    CUDA_CHECK(cudaMemcpy(d_b, h_b.data(), elems * sizeof(__half),
                          cudaMemcpyHostToDevice));

    int status = EXIT_SUCCESS;

    // Phase 1: one tile, and the question of who holds what. The probe
    // tiles hold one unique value per cell, `row * 8 + col` for the 16x8
    // A and `200 + row * 8 + col` for the 8x8 B, so a value read back out
    // of a register names the cell it came from. Every value is an
    // integer below 2048 and therefore exact in __half, and no dot
    // product here reaches 2^24, so the tile check stays exact too.
    std::vector<float> h_want(elems);
    {
        std::vector<__half> h_a16(16 * 8);
        std::vector<float> h_a16f(16 * 8);
        for (int i = 0; i < 16 * 8; ++i) {
            h_a16f[i] = static_cast<float>(i);
            h_a16[i] = __float2half(h_a16f[i]);
        }
        std::vector<__half> h_b8(8 * 8);
        std::vector<float> h_b8f(8 * 8);
        for (int i = 0; i < 8 * 8; ++i) {
            h_b8f[i] = static_cast<float>(200 + i);
            h_b8[i] = __float2half(h_b8f[i]);
        }
        __half* d_a16 = nullptr;
        __half* d_b8 = nullptr;
        CUDA_CHECK(cudaMalloc(&d_a16, 16 * 8 * sizeof(__half)));
        CUDA_CHECK(cudaMalloc(&d_b8, 8 * 8 * sizeof(__half)));
        CUDA_CHECK(cudaMemcpy(d_a16, h_a16.data(), 16 * 8 * sizeof(__half),
                              cudaMemcpyHostToDevice));
        CUDA_CHECK(cudaMemcpy(d_b8, h_b8.data(), 8 * 8 * sizeof(__half),
                              cudaMemcpyHostToDevice));

        probeFragmentMap<<<1, 32>>>(d_a16, d_b8, d_c, d_aFrag, d_bFrag);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());

        std::vector<float> h_d(16 * 8);
        std::vector<float> h_aFrag(32 * 4);
        std::vector<float> h_bFrag(32 * 2);
        CUDA_CHECK(cudaMemcpy(h_d.data(), d_c, 16 * 8 * sizeof(float),
                              cudaMemcpyDeviceToHost));
        CUDA_CHECK(cudaMemcpy(h_aFrag.data(), d_aFrag, 32 * 4 * sizeof(float),
                              cudaMemcpyDeviceToHost));
        CUDA_CHECK(cudaMemcpy(h_bFrag.data(), d_bFrag, 32 * 2 * sizeof(float),
                              cudaMemcpyDeviceToHost));

        matmulCpu(h_a16f.data(), h_b8f.data(), h_want.data(), 16, 8, 8);
        const size_t bad = firstMismatch(h_d.data(), h_want.data(), 16 * 8);
        if (bad != 16 * 8) {
            std::fprintf(stderr,
                         "m16n8k8 tile wrong at row %zu col %zu: got %.9g, "
                         "want %.9g\n",
                         bad / 8, bad % 8, h_d[bad], h_want[bad]);
            status = EXIT_FAILURE;
        } else {
            std::printf(
                "m16n8k8 tile: exact match against the double "
                "reference\n");
        }

        // What ldmatrix actually delivered, decoded from the values, next
        // to what the PTX fragment tables say it should have. The tables
        // are the prediction; the registers are the observation.
        std::printf(
            "A and B fragment ownership, decoded from the "
            "registers (row,col):\n");
        int ownBad = 0;
        for (int l = 0; l < 32; ++l) {
            const int group = l / 4;
            const int member = l % 4;
            int row[6];
            int col[6];
            for (int i = 0; i < 4; ++i) {
                decodeCell(h_aFrag[l * 4 + i], 0.0f, 8, 16 * 8, &row[i],
                           &col[i]);
                if (row[i] != group + ((i < 2) ? 0 : 8) ||
                    col[i] != member * 2 + (i & 1)) {
                    ++ownBad;
                }
            }
            for (int i = 0; i < 2; ++i) {
                decodeCell(h_bFrag[l * 2 + i], 200.0f, 8, 8 * 8, &row[4 + i],
                           &col[4 + i]);
                if (row[4 + i] != member * 2 + i || col[4 + i] != group) {
                    ++ownBad;
                }
            }
            std::printf(
                "lane %2d  a: (%2d,%2d) (%2d,%2d) (%2d,%2d) (%2d,%2d)"
                "  b: (%2d,%2d) (%2d,%2d)\n",
                l, row[0], col[0], row[1], col[1], row[2], col[2], row[3],
                col[3], row[4], col[4], row[5], col[5]);
        }
        if (ownBad != 0) {
            std::fprintf(stderr,
                         "%d of 192 fragment elements sit where the PTX "
                         "tables do not predict\n",
                         ownBad);
            status = EXIT_FAILURE;
        } else {
            std::printf(
                "fragment ownership: all 192 elements match the PTX "
                "tables\n");
        }

        // Phase 1b: the sm_80 shape, on hardware that has it.
        if (prop.major >= 8) {
            std::vector<__half> h_a256(16 * 16);
            std::vector<float> h_a256f(16 * 16);
            for (int row = 0; row < 16; ++row) {
                for (int col = 0; col < 16; ++col) {
                    h_a256[row * 16 + col] = h_a[row * kSize + col];
                    h_a256f[row * 16 + col] = h_af[row * kSize + col];
                }
            }
            std::vector<__half> h_b128(16 * 8);
            std::vector<float> h_b128f(16 * 8);
            for (int row = 0; row < 16; ++row) {
                for (int col = 0; col < 8; ++col) {
                    h_b128[row * 8 + col] = h_b[row * kSize + col];
                    h_b128f[row * 8 + col] = h_bf[row * kSize + col];
                }
            }
            __half* d_a256 = nullptr;
            __half* d_b128 = nullptr;
            CUDA_CHECK(cudaMalloc(&d_a256, 16 * 16 * sizeof(__half)));
            CUDA_CHECK(cudaMalloc(&d_b128, 16 * 8 * sizeof(__half)));
            CUDA_CHECK(cudaMemcpy(d_a256, h_a256.data(),
                                  16 * 16 * sizeof(__half),
                                  cudaMemcpyHostToDevice));
            CUDA_CHECK(cudaMemcpy(d_b128, h_b128.data(),
                                  16 * 8 * sizeof(__half),
                                  cudaMemcpyHostToDevice));

            probeTileK16<<<1, 32>>>(d_a256, d_b128, d_c);
            CUDA_CHECK(cudaGetLastError());
            CUDA_CHECK(cudaDeviceSynchronize());
            CUDA_CHECK(cudaMemcpy(h_d.data(), d_c, 16 * 8 * sizeof(float),
                                  cudaMemcpyDeviceToHost));
            matmulCpu(h_a256f.data(), h_b128f.data(), h_want.data(), 16, 8, 16);
            const size_t bad16 =
                firstMismatch(h_d.data(), h_want.data(), 16 * 8);
            if (bad16 != 16 * 8) {
                std::fprintf(stderr,
                             "m16n8k16 tile wrong at row %zu col %zu: got "
                             "%.9g, want %.9g\n",
                             bad16 / 8, bad16 % 8, h_d[bad16], h_want[bad16]);
                status = EXIT_FAILURE;
            } else {
                std::printf(
                    "m16n8k16 tile: exact match against the double "
                    "reference\n");
            }
            CUDA_CHECK(cudaFree(d_a256));
            CUDA_CHECK(cudaFree(d_b128));
        } else {
            std::printf(
                "m16n8k16 tile: skipped, needs sm_80 or newer "
                "(this GPU is sm_%d%d)\n",
                prop.major, prop.minor);
        }

        CUDA_CHECK(cudaFree(d_a16));
        CUDA_CHECK(cudaFree(d_b8));
    }

    // Phase 2: the full matmul, checked before it is timed.
    const dim3 grid(kSize / kBlockN, kSize / kBlockM);
    if (status == EXIT_SUCCESS) {
        matmulMmaSync<<<grid, kThreadsPerBlock>>>(d_a, d_b, d_c, kSize);
        CUDA_CHECK(cudaGetLastError());
        CUDA_CHECK(cudaDeviceSynchronize());

        std::vector<float> h_c(elems);
        CUDA_CHECK(cudaMemcpy(h_c.data(), d_c, elems * sizeof(float),
                              cudaMemcpyDeviceToHost));
        matmulCpu(h_af.data(), h_bf.data(), h_want.data(), kSize, kSize, kSize);
        const size_t bad = firstMismatch(h_c.data(), h_want.data(), elems);
        if (bad != elems) {
            std::fprintf(stderr,
                         "matmulMmaSync wrong at row %zu col %zu: got %.9g, "
                         "want %.9g\n",
                         bad / kSize, bad % kSize, h_c[bad], h_want[bad]);
            status = EXIT_FAILURE;
        } else {
            std::printf(
                "matmulMmaSync at n = %d: exact match against the "
                "double reference\n",
                kSize);
        }
    }

    if (status == EXIT_SUCCESS) {
        const float ms = timeKernel([&] {
            matmulMmaSync<<<grid, kThreadsPerBlock>>>(d_a, d_b, d_c, kSize);
        });
        const double flops = 2.0 * kSize * static_cast<double>(kSize) * kSize;
        std::printf(
            "matmulMmaSync: %.3f ms, %.1f GFLOP/s (FP16 in, FP32 "
            "accumulate, mean of %d runs, copies not included)\n",
            ms, flops / (static_cast<double>(ms) * 1.0e-3) / 1.0e9, kTimedRuns);
    }

    CUDA_CHECK(cudaFree(d_a));
    CUDA_CHECK(cudaFree(d_b));
    CUDA_CHECK(cudaFree(d_c));
    CUDA_CHECK(cudaFree(d_aFrag));
    CUDA_CHECK(cudaFree(d_bFrag));
    return status;
}