From 562a6cbf6a1d66bc22997ec2285f2ffe61471065 Mon Sep 17 00:00:00 2001 From: mcencini Date: Wed, 7 Oct 2026 13:17:56 +0200 Subject: [PATCH] Tile the PERK kernels as matrix products The ported PERK kernels formed each voxel's features one thread at a time, reading every frequency from global memory once per multiply. They ran 2.5x slower than the Triton kernels they replaced. Each program now stages a block of signals and frequencies in shared memory, and each thread forms two voxels by a block of features in registers, reading weights four at a time. The accumulators are stored through unrolled loops, so they stay in registers and nothing spills. On an RTX 4060 (sm_89), 200k voxels, 24 contrasts, 256 features: forward 2.4 ms against 5.1 ms for Triton and 12.9 ms before; adjoint 4.6 ms against 5.4 ms and 14.0 ms. Below about 4k voxels a program's serial chain over the features is longer than Triton's, by up to 45 us. On the host the same source runs a program's threads one after another between barriers, so the host build still checks the card's code. The parity test now runs on the card too. Co-Authored-By: Claude Opus 5.5 --- src/blochsim/_kernels.hpp | 4 +- src/blochsim/_perk_kernels.hpp | 390 ++++++++++++++++++++++----- src/blochsim/estimators/_perk_gpu.py | 8 +- tests/estimators/test_perk_kernel.py | 21 +- 4 files changed, 339 insertions(+), 84 deletions(-) diff --git a/src/blochsim/_kernels.hpp b/src/blochsim/_kernels.hpp index 50ce00f1..b76cca32 100644 --- a/src/blochsim/_kernels.hpp +++ b/src/blochsim/_kernels.hpp @@ -819,8 +819,8 @@ inline constexpr KernelInfo KERNELS[] = { {"_epg_jvp_kernel", "t1,t2,m0,b1,b1_phase,b0,inversion_efficiency,diffusion,velocity,bound_fraction,exchange_rate,t1_bound,pool_b_fraction,pool_b_exchange,t1_pool_b,t2_pool_b,pool_b_shift,duration,kind,flip,phase,action,output_index,shim_index,tangent_t1,tangent_t2,tangent_m0,tangent_b1,tangent_b1_phase,tangent_b0,tangent_inversion_efficiency,tangent_diffusion,tangent_velocity,tangent_bound_fraction,tangent_exchange_rate,tangent_t1_bound,tangent_pool_b_fraction,tangent_pool_b_exchange,tangent_t1_pool_b,tangent_t2_pool_b,tangent_pool_b_shift,tangent_duration,tangent_flip,tangent_phase,saturation,rf_frequency,profile,profile_index,lineshape,pairs,pair_index,pair_direction,duration_row,pool_table,output_real,output_imag,atom_count,train_count,event_count,output_count,flow_scale,washout_scale,profile_step,lineshape_step,state_count,single_train,atom_stride,shim_rows,shimmed,locations,profiled,profile_bins,dynamic,broadened,lineshape_bins,pools,narrow,tabulated,off_axis,moving,diffusing,transmit,density,inverting,block_states,problems", "ppppppppppppppppppppppppppppppppppppppppppppppppppppppppiiiiffffiiiiiiiiiiiiiiiiiiiiii", 84, 85, -1}, {"_pooled_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,output_real,output_imag,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,keep,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiii", 70, 69, 69}, {"_pooled_adjoint_kernel", "m0,b1,b1_phase,b0,efficiency,diffusion,velocity,dm0,db1,db1_phase,db0,defficiency,ddiffusion,dvelocity,duration,kind,flip,phase,action,output_index,shim_index,saturation,rf_frequency,dduration,dflip,dphase,table,dtable,pool_index,profile,profile_index,lineshape,pairs,pair_index,dpairs,grad_real,grad_imag,grad_tissue,dgrad_tissue,grad_duration,dgrad_duration,grad_flip,dgrad_flip,grad_phase,dgrad_phase,grad_table,dgrad_table,grad_pairs,dgrad_pairs,trajectory,base,atom_count,event_count,output_count,state_count,rows,flow_scale,washout_scale,profile_step,lineshape_step,locations,profile_bins,lineshape_bins,m0_row,b1_row,b1_phase_row,b0_row,efficiency_row,diffusion_row,velocity_row,n,m,blocks,planes,atom_stride,shimmed,profiled,dynamic,directed_pairs,directed_table,following,off_axis,moving,diffusing,transmit,density,inverting,P,S", "ppppppppppppppppppppppppppppppppppppppppppppppppppiiiiiiffffiiiiiiiiiiiiiiiiiiiiiiiiiiiii", 88, 87, 87}, - {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "pppppppiiiifi", 12, -1, -1}, - {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,BLOCK_VOXELS", "ppppppiiiifi", 11, -1, -1}, + {"_regress_kernel", "signal,frequency,phase,feature_mean,weight,parameter_mean,output,voxels,contrasts,features,parameters,scale,threads", "pppppppiiiifi", 12, -1, -1}, + {"_regress_vjp_kernel", "signal,frequency,phase,weight,cotangent,output,voxels,contrasts,features,parameters,scale,threads", "ppppppiiiifi", 11, -1, -1}, }; #define BLOCHSIM_FOR_EACH_KERNEL(X) \ diff --git a/src/blochsim/_perk_kernels.hpp b/src/blochsim/_perk_kernels.hpp index b77ffb4b..d03622a6 100644 --- a/src/blochsim/_perk_kernels.hpp +++ b/src/blochsim/_perk_kernels.hpp @@ -1,112 +1,358 @@ -// The PERK feature map and its regression, fused, one voxel per thread. +// The PERK feature map and its regression, fused. // // y = parameter_mean + (scale * cos(W @ x + b) - feature_mean) @ weight.T // -// A block of features is formed and consumed into the output accumulator in -// registers, so the ``(voxels, features)`` matrix never exists. The adjoint -// does the same and forms the angles again rather than keeping them. Every -// array is contiguous and row-major. +// Two matrix products with a cosine between, tiled the way a matrix product +// is: a program stages a block of signals and a block of frequencies in shared +// memory, and each thread forms the angles of VOXELS voxels by a block of +// features in registers, so every value read from shared memory feeds several +// multiplies. The cosine and the second product consume the block while it is +// still in registers, so the ``(voxels, features)`` matrix never exists. The +// adjoint forms the angles again rather than keeping them. Every array is +// contiguous and row-major. +// +// Staged blocks are zero past the last feature, parameter or voxel, and a zero +// weight is what removes a padded feature from both products, so the inner +// loops carry no bounds. On a card a program is THREADS threads; on the host +// the same source runs them one after another between barriers. + +// Threads per program, and the voxels each holds. ``_perk_gpu`` launches +// programs of ``THREADS`` threads over ``THREADS * VOXELS`` voxels. +constexpr int THREADS = 64; +constexpr int VOXELS = 2; +constexpr int BLOCK_VOXELS = THREADS * VOXELS; +// Contrasts staged at once. +constexpr int CONTRASTS = 32; +// Features formed at once by the forward pass, and parameters it accumulates. +constexpr int FEATURES = 32; +constexpr int PARAMETERS = 8; +// Features formed at once by the adjoint, and contrasts of the gradient it +// accumulates: it holds the angles, their cotangents and the gradient at once. +constexpr int ADJOINT_FEATURES = 16; +constexpr int GRADIENT = 32; +// A row of staged weights is read four at a time, so it is padded to keep each +// row 16-byte aligned and the rows on different banks. +constexpr int PAD = 4; + +#if defined(BLOCHSIM_SIMT) +#define PERK_SHARED __shared__ +#define PERK_SHARED_ROWS __shared__ __align__(16) +#define PERK_SYNC() __syncthreads() +// The body runs once, as this thread. +#define PERK_EACH_THREAD(t) \ + for (int t = static_cast(threadIdx.x), t##_once = 1; t##_once; t##_once = 0) +// What a thread keeps across a barrier: its own registers. +template +struct Own { + T value; + BSK_HD T& operator[](int) { return value; } +}; +#else +#define PERK_SHARED static thread_local +#define PERK_SHARED_ROWS alignas(16) static thread_local +#define PERK_SYNC() static_cast(0) +// The body runs as each thread in turn, so a barrier is the end of the loop. +#define PERK_EACH_THREAD(t) for (int t = 0; t < THREADS; ++t) +template +struct Own { + T value[THREADS]; + T& operator[](int t) { return value[t]; } +}; +#endif -// Features formed at once, which is how often a voxel's signal is read. -constexpr int FEATURE_BLOCK = 32; -// Parameters accumulated at once by the forward pass. -constexpr int PARAMETER_BLOCK = 16; -// Contrasts accumulated at once by the adjoint. -constexpr int CONTRAST_BLOCK = 32; +struct alignas(16) Four { + float x, y, z, w; +}; -// The angles ``W @ x + b`` of features ``first`` onward, as many as there are. -template -BSK_HD void _angles(const float* signal, const float* frequency, const float* phase, - const Voxel& voxel, const Live& live, std::int64_t contrasts, - std::int64_t features, std::int64_t first, bsk::V* angle) { - for (int j = 0; j < FEATURE_BLOCK; ++j) { - angle[j] = bsk::V(first + j < features ? phase[first + j] : 0.0f); +BSK_HD const Four& four(const float* address) { + return *reinterpret_cast(address); +} + +BSK_HD int clamp_width(std::int64_t left, int block) { + return left < block ? static_cast(left) : block; +} + +// Signals of voxels ``first`` onward, contrasts ``c0`` onward, ``width`` of +// them, into ``staged[contrast][voxel]``. Consecutive threads copy consecutive +// elements of the rows, which are contiguous in memory. +BSK_HD void _stage_signals(const float* signal, std::int64_t voxels, std::int64_t contrasts, + std::int64_t first, std::int64_t c0, int width, + float (*staged)[BLOCK_VOXELS + 1], int t) { + for (int i = t; i < BLOCK_VOXELS * width; i += THREADS) { + const int v = i / width; + const int c = i - v * width; + const std::int64_t voxel = first + v; + staged[c][v] = voxel < voxels ? signal[voxel * contrasts + c0 + c] : 0.0f; } - for (std::int64_t contrast = 0; contrast < contrasts; ++contrast) { - const auto value = bsk::ld(signal + voxel * contrasts + contrast, live, 0.0f); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j < features) { - angle[j] = angle[j] + value * frequency[(first + j) * contrasts + contrast]; +} + +// Frequencies of ``width`` features ``f0`` onward against contrasts ``c0`` +// onward, into ``staged[contrast][feature]``, zero past the last feature. +template +BSK_HD void _stage_frequencies(const float* frequency, std::int64_t contrasts, + std::int64_t f0, int features, std::int64_t c0, int width, + float (*staged)[BLOCK + PAD], int t) { + for (int i = t; i < BLOCK * width; i += THREADS) { + const int j = i / width; + const int c = i - j * width; + staged[c][j] = j < features ? frequency[(f0 + j) * contrasts + c0 + c] : 0.0f; + } +} + +// The angles ``W @ x + b`` of a thread's voxels over a block of features, +// accumulated over the staged contrasts. +template +BSK_HD void _accumulate_angles(const float (*signals)[BLOCK_VOXELS + 1], + const float (*frequencies)[BLOCK + PAD], int width, + float (&angle)[VOXELS][BLOCK], int t) { + for (int c = 0; c < width; ++c) { + float x[VOXELS]; +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + x[r] = signals[c][t + r * THREADS]; + } +#pragma unroll + for (int j = 0; j < BLOCK; j += 4) { + const Four w = four(&frequencies[c][j]); +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + angle[r][j] += x[r] * w.x; + angle[r][j + 1] += x[r] * w.y; + angle[r][j + 2] += x[r] * w.z; + angle[r][j + 3] += x[r] * w.w; + } + } + } +} + +// The angles of every voxel of the program over ``width`` features ``f0`` +// onward, the phases already staged in ``offset``. Ends at a barrier. +template +BSK_HD void _angles(const float* signal, const float* frequency, std::int64_t voxels, + std::int64_t contrasts, std::int64_t first, std::int64_t f0, int width, + const float* offset, float (*signals)[BLOCK_VOXELS + 1], + float (*frequencies)[BLOCK + PAD], Own& angle) { + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { + angle[t][r][j] = offset[j]; } } } + for (std::int64_t c0 = 0; c0 < contrasts; c0 += CONTRASTS) { + const int staged = clamp_width(contrasts - c0, CONTRASTS); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + _stage_signals(signal, voxels, contrasts, first, c0, staged, signals, t); + _stage_frequencies(frequency, contrasts, f0, width, c0, staged, frequencies, + t); + } + PERK_SYNC(); + PERK_EACH_THREAD(t) { + _accumulate_angles(signals, frequencies, staged, angle[t], t); + } + } + PERK_SYNC(); } -// One block of voxels, from signal to parameters. +// One program of voxels, from signal to parameters. BSK_HD void _regress_kernel(const float* signal, const float* frequency, const float* phase, const float* feature_mean, const float* weight, const float* parameter_mean, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, - std::int64_t parameters, float scale, std::int64_t BLOCK_VOXELS) { - const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); - const auto live = voxel < voxels; - bsk::V angle[FEATURE_BLOCK]; - for (std::int64_t base = 0; base < parameters; base += PARAMETER_BLOCK) { - bsk::V total[PARAMETER_BLOCK]; - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - total[k] = bsk::V(0.0f); + std::int64_t parameters, float scale, std::int64_t threads) { + static_cast(threads); + PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; + PERK_SHARED_ROWS float frequencies[CONTRASTS][FEATURES + PAD]; + PERK_SHARED_ROWS float weights[FEATURES][PARAMETERS]; + PERK_SHARED float offset[FEATURES]; + PERK_SHARED float mean[FEATURES]; + const std::int64_t first = bsk::program_id(0) * BLOCK_VOXELS; + Own angle; + Own total; + for (std::int64_t p0 = 0; p0 < parameters; p0 += PARAMETERS) { + const int held = clamp_width(parameters - p0, PARAMETERS); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int k = 0; k < PARAMETERS; ++k) { + total[t][r][k] = 0.0f; + } + } } - for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { - _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j >= features) { - break; + for (std::int64_t f0 = 0; f0 < features; f0 += FEATURES) { + const int width = clamp_width(features - f0, FEATURES); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int j = t; j < FEATURES; j += THREADS) { + offset[j] = j < width ? phase[f0 + j] : 0.0f; + mean[j] = j < width ? feature_mean[f0 + j] : 0.0f; + } + for (int i = t; i < FEATURES * PARAMETERS; i += THREADS) { + const int j = i / PARAMETERS; + const int k = i - j * PARAMETERS; + weights[j][k] = + j < width && k < held ? weight[(p0 + k) * features + f0 + j] : 0.0f; } - const auto mapped = scale * bsk::cos(angle[j]) - feature_mean[first + j]; - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - if (base + k < parameters) { - total[k] = total[k] + mapped * weight[(base + k) * features + first + j]; + } + PERK_SYNC(); + _angles(signal, frequency, voxels, contrasts, first, f0, width, offset, + signals, frequencies, angle); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int j = 0; j < FEATURES; ++j) { + const float m = mean[j]; +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const float mapped = scale * cosf(angle[t][r][j]) - m; +#pragma unroll + for (int k = 0; k < PARAMETERS; k += 4) { + const Four w = four(&weights[j][k]); + total[t][r][k] += mapped * w.x; + total[t][r][k + 1] += mapped * w.y; + total[t][r][k + 2] += mapped * w.z; + total[t][r][k + 3] += mapped * w.w; + } } } } } - for (int k = 0; k < PARAMETER_BLOCK; ++k) { - if (base + k < parameters) { - bsk::st(output + voxel * parameters + base + k, - total[k] + parameter_mean[base + k], live); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; +#pragma unroll + for (int k = 0; k < PARAMETERS; ++k) { + if (voxel < voxels && k < held) { + output[voxel * parameters + p0 + k] = + total[t][r][k] + parameter_mean[p0 + k]; + } + } } } } } -// The derivative of one block of voxels with respect to their signals. +// The derivative of one program of voxels with respect to their signals. BSK_HD void _regress_vjp_kernel(const float* signal, const float* frequency, const float* phase, const float* weight, const float* cotangent, float* output, std::int64_t voxels, std::int64_t contrasts, std::int64_t features, std::int64_t parameters, float scale, - std::int64_t BLOCK_VOXELS) { - const auto voxel = bsk::program_id(0) * BLOCK_VOXELS + bsk::arange_x(); - const auto live = voxel < voxels; - bsk::V angle[FEATURE_BLOCK]; - for (std::int64_t base = 0; base < contrasts; base += CONTRAST_BLOCK) { - bsk::V gradient[CONTRAST_BLOCK]; - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - gradient[c] = bsk::V(0.0f); + std::int64_t threads) { + static_cast(threads); + constexpr int BLOCK = ADJOINT_FEATURES; + PERK_SHARED float signals[CONTRASTS][BLOCK_VOXELS + 1]; + PERK_SHARED_ROWS float frequencies[CONTRASTS][BLOCK + PAD]; + PERK_SHARED_ROWS float back[BLOCK][GRADIENT + PAD]; + PERK_SHARED_ROWS float weights[PARAMETERS][BLOCK]; + PERK_SHARED float offset[BLOCK]; + const std::int64_t first = bsk::program_id(0) * BLOCK_VOXELS; + Own angle; + Own through; + Own gradient; + for (std::int64_t g0 = 0; g0 < contrasts; g0 += GRADIENT) { + const int held = clamp_width(contrasts - g0, GRADIENT); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int c = 0; c < GRADIENT; ++c) { + gradient[t][r][c] = 0.0f; + } + } } - for (std::int64_t first = 0; first < features; first += FEATURE_BLOCK) { - _angles(signal, frequency, phase, voxel, live, contrasts, features, first, angle); - for (int j = 0; j < FEATURE_BLOCK; ++j) { - if (first + j >= features) { - break; + for (std::int64_t f0 = 0; f0 < features; f0 += BLOCK) { + const int width = clamp_width(features - f0, BLOCK); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int j = t; j < BLOCK; j += THREADS) { + offset[j] = j < width ? phase[f0 + j] : 0.0f; + } + for (int i = t; i < BLOCK * held; i += THREADS) { + const int j = i / held; + const int c = i - j * held; + back[j][c] = j < width ? frequency[(f0 + j) * contrasts + g0 + c] : 0.0f; + } + for (int i = t; i < BLOCK * (GRADIENT - held); i += THREADS) { + const int j = i / (GRADIENT - held); + back[j][held + i - j * (GRADIENT - held)] = 0.0f; } - bsk::V through(0.0f); - for (std::int64_t p = 0; p < parameters; ++p) { - through = through - + bsk::ld(cotangent + voxel * parameters + p, live, 0.0f) - * weight[p * features + first + j]; + } + PERK_SYNC(); + _angles(signal, frequency, voxels, contrasts, first, f0, width, offset, + signals, frequencies, angle); + // What reaches each feature from the parameters: cotangent @ weight. + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { + through[t][r][j] = 0.0f; + } + } + } + for (std::int64_t p0 = 0; p0 < parameters; p0 += PARAMETERS) { + const int count = clamp_width(parameters - p0, PARAMETERS); + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int i = t; i < PARAMETERS * BLOCK; i += THREADS) { + const int k = i / BLOCK; + const int j = i - k * BLOCK; + weights[k][j] = + k < count && j < width ? weight[(p0 + k) * features + f0 + j] : 0.0f; + } } - through = through * (-scale) * bsk::sin(angle[j]); - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - if (base + c < contrasts) { - gradient[c] = gradient[c] - + through * frequency[(first + j) * contrasts + base + c]; + PERK_SYNC(); + PERK_EACH_THREAD(t) { + for (int k = 0; k < count; ++k) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; + const float pulled = + voxel < voxels ? cotangent[voxel * parameters + p0 + k] : 0.0f; +#pragma unroll + for (int j = 0; j < BLOCK; j += 4) { + const Four w = four(&weights[k][j]); + through[t][r][j] += pulled * w.x; + through[t][r][j + 1] += pulled * w.y; + through[t][r][j + 2] += pulled * w.z; + through[t][r][j + 3] += pulled * w.w; + } + } + } + } + } + PERK_EACH_THREAD(t) { +#pragma unroll + for (int j = 0; j < BLOCK; ++j) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const float slope = -scale * sinf(angle[t][r][j]) * through[t][r][j]; +#pragma unroll + for (int c = 0; c < GRADIENT; c += 4) { + const Four w = four(&back[j][c]); + gradient[t][r][c] += slope * w.x; + gradient[t][r][c + 1] += slope * w.y; + gradient[t][r][c + 2] += slope * w.z; + gradient[t][r][c + 3] += slope * w.w; + } } } } } - for (int c = 0; c < CONTRAST_BLOCK; ++c) { - if (base + c < contrasts) { - bsk::st(output + voxel * contrasts + base + c, gradient[c], live); + PERK_EACH_THREAD(t) { +#pragma unroll + for (int r = 0; r < VOXELS; ++r) { + const std::int64_t voxel = first + t + r * THREADS; +#pragma unroll + for (int c = 0; c < GRADIENT; ++c) { + if (voxel < voxels && c < held) { + output[voxel * contrasts + g0 + c] = gradient[t][r][c]; + } + } } } } diff --git a/src/blochsim/estimators/_perk_gpu.py b/src/blochsim/estimators/_perk_gpu.py index b68088b7..909ace13 100644 --- a/src/blochsim/estimators/_perk_gpu.py +++ b/src/blochsim/estimators/_perk_gpu.py @@ -23,7 +23,9 @@ from .._gpu_launch import Kernel, cdiv -#: Voxels per program, one per thread. +#: Threads per program, and the voxels a program holds: ``THREADS`` and +#: ``THREADS * VOXELS`` in ``_perk_kernels.hpp``. +_THREADS = 64 _BLOCK_VOXELS = 128 _regress_kernel = Kernel("_regress_kernel") @@ -71,7 +73,7 @@ def regress( features, parameters, math.sqrt(2.0 / features), - _BLOCK_VOXELS, + _THREADS, ) return output @@ -108,6 +110,6 @@ def regress_vjp( features, parameters, math.sqrt(2.0 / features), - _BLOCK_VOXELS, + _THREADS, ) return output diff --git a/tests/estimators/test_perk_kernel.py b/tests/estimators/test_perk_kernel.py index fee8e26e..1f160011 100644 --- a/tests/estimators/test_perk_kernel.py +++ b/tests/estimators/test_perk_kernel.py @@ -207,14 +207,18 @@ def test_a_single_voxel_goes_through_a_kernel_tiled_for_many(device) -> None: assert estimator(measured[:1]).shape == (1, 2) -def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint() -> None: - """The kernels the card runs, compiled for the host, against Torch. +@pytest.mark.parametrize("device", DEVICES) +def test_the_gpu_kernels_are_the_fused_line_and_its_adjoint(device) -> None: + """The kernels the card runs, on the card and compiled for the host, + against Torch. Feature, parameter and contrast counts that are not multiples of the - blocks the kernels walk them in, so every edge of the tiling is read. + blocks the kernels walk them in, and more voxels than one program holds, + so every edge of the tiling is read. """ gpu = pytest.importorskip("blochsim.estimators._perk_gpu") - pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") + if device == "cpu": + pytest.importorskip("blochsim._gpu_host", reason="the host build is Linux only") generator = torch.Generator().manual_seed(0) voxels, contrasts, features, parameters = 300, 37, 45, 17 signals = torch.randn(voxels, contrasts, generator=generator) @@ -238,10 +242,13 @@ def line(x: torch.Tensor) -> torch.Tensor: expected = line(x) (expected_gradient,) = torch.autograd.grad(expected, x, cotangent.double()) + def on(*tensors: torch.Tensor) -> list[torch.Tensor]: + return [tensor.float().to(device) for tensor in tensors] + estimated = gpu.regress( - signals, frequency, phase, feature_mean, weight, parameter_mean - ) - gradient = gpu.regress_vjp(cotangent, signals, frequency, phase, weight) + *on(signals, frequency, phase, feature_mean, weight, parameter_mean) + ).cpu() + gradient = gpu.regress_vjp(*on(cotangent, signals, frequency, phase, weight)).cpu() torch.testing.assert_close( estimated.double(), expected.detach(), atol=1e-5, rtol=1e-5