perf(cuda): Wave 3 — TMA bulk async tile loads for Hopper (sm_90+)

Add cp.async.bulk.shared::cta.global tile loads guarded by
#if __CUDA_ARCH__ >= 900 in common_device_functions.cuh:

- cooperative_load_tile_tma(): thread-0-only bulk copy via inline PTX,
  freeing 31 warp threads for compute overlap. 16KB chunks with
  commit_group/wait_group barrier.
- cooperative_load_tile_float4(): renamed original for fallback.
- cooperative_load_tile(): dispatch wrapper (compile-time selection).

Architecture-aware NVRTC compilation (compile_ptx_for_device):
- Queries GPU compute capability, passes -arch=compute_XX to NVRTC.
- Enables __CUDA_ARCH__ in kernels so TMA guard activates on Hopper.
- Wired into all 3 runtime compilation sites (experience collector,
  backtest evaluator, PPO collector).

Also fixes pre-existing clippy: vh * 1 identity op in weight estimate.

79 cuda_pipeline tests pass, 0 clippy errors.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-11 22:36:43 +01:00
parent 789dc86589
commit 789fb50dfe
5 changed files with 173 additions and 18 deletions

View File

@@ -172,11 +172,13 @@ __device__ void noisy_matvec_leaky_relu(
#endif
/**
* Cooperatively load a weight tile from global to shared memory.
* Cooperatively load a weight tile from global to shared memory (float4 path).
* All threads in the block participate. Must call __syncthreads() after.
* Uses float4 vectorized loads where possible for 4x bandwidth.
*
* This is the pre-Hopper fallback used on sm_70/sm_80/sm_86/sm_89.
*/
__device__ void cooperative_load_tile(
__device__ __forceinline__ void cooperative_load_tile_float4(
float* __restrict__ shmem_dst,
const float* __restrict__ global_src,
int count /* total floats to load */
@@ -194,6 +196,100 @@ __device__ void cooperative_load_tile(
}
}
#if __CUDA_ARCH__ >= 900
/**
* TMA-accelerated tile load from global to shared memory (Hopper sm_90+).
*
* Uses cp.async.bulk.shared::cta.global to perform a bulk asynchronous copy
* from global (HBM) to shared memory via the Tensor Memory Accelerator.
* Only thread 0 issues the copy instruction; all other threads are freed
* for compute overlap. The caller's __syncthreads() ensures completion
* visibility (TMA commits to the async group, wait_group 0 drains it).
*
* cp.async.bulk transfers up to 128 bytes per instruction in the base form,
* so we issue ceil(byte_count / 128) sequential bulk copies from thread 0.
* This still wins over the float4 scatter because:
* (a) TMA bypasses L1/L2 pollution — direct HBM→SMEM path
* (b) No warp divergence — 31 threads skip straight to compute
* (c) Overlaps with any independent arithmetic on the other warps
*
* Fallback: tiles larger than practical TMA dispatch (>64KB, rare) fall
* through to the float4 cooperative path to avoid excessive PTX overhead.
*
* Constraints:
* - src must be 16-byte aligned (guaranteed by cudarc device alloc)
* - shmem_dst must be 16-byte aligned (guaranteed by __shared__ layout)
* - count must be > 0
* - Caller MUST issue __syncthreads() (or mbarrier) after return
*/
__device__ __forceinline__ void cooperative_load_tile_tma(
float* __restrict__ shmem_dst,
const float* __restrict__ global_src,
int count /* total floats to load */
) {
/* cp.async.bulk max per-instruction is limited; we chunk at 16KB
* (4096 floats) which is well within the hardware capability and
* keeps PTX emission tractable. Each chunk is a single bulk copy. */
const int CHUNK_FLOATS = 4096; /* 16 KB per chunk */
const int byte_count = count * (int)sizeof(float);
/* Fallback for very large tiles (>256KB): should not happen with
* SHMEM_TILE_ROWS=64 × SHMEM_MAX_IN_DIM=256 = 64KB, but guard. */
if (byte_count > 256 * 1024) {
cooperative_load_tile_float4(shmem_dst, global_src, count);
return;
}
if (threadIdx.x == 0) {
/* Issue bulk async copies in chunks.
* cp.async.bulk.shared::cta.global [dst], [src], size;
* size operand is in bytes, immediate or register. */
int remaining = count;
int offset = 0;
while (remaining > 0) {
int chunk = remaining < CHUNK_FLOATS ? remaining : CHUNK_FLOATS;
int chunk_bytes = chunk * (int)sizeof(float);
/* Inline PTX: cp.async.bulk.shared::cta.global
* Operands: .dst (shared), .src (global), .size (bytes) */
asm volatile(
"cp.async.bulk.shared::cta.global [%0], [%1], %2;\n"
:
: "r"((unsigned int)__cvta_generic_to_shared(shmem_dst + offset)),
"l"((unsigned long long)(global_src + offset)),
"r"((unsigned int)chunk_bytes)
: "memory"
);
offset += chunk;
remaining -= chunk;
}
/* Commit all issued bulk copies to an async group and wait. */
asm volatile("cp.async.bulk.commit_group;\n" ::: "memory");
asm volatile("cp.async.bulk.wait_group 0;\n" ::: "memory");
}
/* Caller issues __syncthreads() to broadcast completion to all threads.
* This is already the case in every TILE_LAYER_* macro and inline site. */
}
#endif /* __CUDA_ARCH__ >= 900 */
/**
* Cooperatively load a weight tile from global to shared memory.
* All threads in the block participate. Must call __syncthreads() after.
*
* Dispatch: on Hopper (sm_90+) uses TMA bulk async copy via thread 0.
* On older architectures uses float4 vectorized cooperative loads.
*/
__device__ __forceinline__ void cooperative_load_tile(
float* __restrict__ shmem_dst,
const float* __restrict__ global_src,
int count /* total floats to load */
) {
#if __CUDA_ARCH__ >= 900
cooperative_load_tile_tma(shmem_dst, global_src, count);
#else
cooperative_load_tile_float4(shmem_dst, global_src, count);
#endif
}
/**
* Matrix-vector multiply from shared memory: output = shmem_W * input + shmem_b.
*

View File

@@ -997,12 +997,15 @@ impl GpuBacktestEvaluator {
shmem_max_in={shmem_max_in_dim} tile_rows={shmem_tile_rows}"
);
let ptx: Ptx = cudarc::nvrtc::compile_ptx(&full_source).map_err(|e| {
MLError::ModelError(format!(
"backtest_forward_kernel CUDA compilation failed: {e}"
))
})?;
let context = self.stream.context();
// Use arch-aware compilation so __CUDA_ARCH__ is defined correctly.
// On Hopper (sm_90+) this enables TMA bulk async tile loads.
let ptx: Ptx = crate::cuda_pipeline::compile_ptx_for_device(&full_source, &context)
.map_err(|e| {
MLError::ModelError(format!(
"backtest_forward_kernel CUDA compilation failed: {e}"
))
})?;
let module = context.load_module(ptx).map_err(|e| {
MLError::ModelError(format!("backtest_forward module load: {e}"))
})?;

View File

@@ -400,12 +400,15 @@ impl GpuExperienceCollector {
shared=[{shared_h1},{shared_h2}] value={value_h} adv={adv_h} atoms_max={num_atoms_max} \
shmem_max_in={shmem_max_in_dim} shmem_tile_rows={shmem_tile_rows}"
);
let ptx: Ptx = cudarc::nvrtc::compile_ptx(&full_source).map_err(|e| {
MLError::ModelError(format!(
"CUDA dqn_experience_kernel compilation failed (is NVRTC available?): {e}"
))
})?;
let context = stream.context();
// Use arch-aware compilation so __CUDA_ARCH__ is defined correctly.
// On Hopper (sm_90+) this enables TMA bulk async tile loads.
let ptx: Ptx = crate::cuda_pipeline::compile_ptx_for_device(&full_source, &context)
.map_err(|e| {
MLError::ModelError(format!(
"CUDA dqn_experience_kernel compilation failed (is NVRTC available?): {e}"
))
})?;
let module = context.load_module(ptx).map_err(|e| {
MLError::ModelError(format!("Failed to load dqn_experience_kernel module: {e}"))
})?;
@@ -572,7 +575,7 @@ impl GpuExperienceCollector {
let dueling_params = state_dim * sh1 + sh1
+ sh1 * sh2 + sh2
+ sh2 * vh + vh
+ vh * 1 + 1 // value out
+ vh + 1 // value out
+ sh2 * ah + ah
+ ah * 5 + 5; // advantage out (5 actions)
// Two dueling networks (online + target)

View File

@@ -196,12 +196,15 @@ impl GpuPpoExperienceCollector {
let common_src = include_str!("common_device_functions.cuh");
let kernel_src = include_str!("ppo_experience_kernel.cu");
let full_source = format!("{}\n{}", common_src, kernel_src);
let ptx: Ptx = cudarc::nvrtc::compile_ptx(&full_source).map_err(|e| {
MLError::ModelError(format!(
"CUDA ppo_experience_kernel compilation failed (is NVRTC available?): {e}"
))
})?;
let context = stream.context();
// Use arch-aware compilation so __CUDA_ARCH__ is defined correctly.
// On Hopper (sm_90+) this enables TMA bulk async tile loads.
let ptx: Ptx = crate::cuda_pipeline::compile_ptx_for_device(&full_source, &context)
.map_err(|e| {
MLError::ModelError(format!(
"CUDA ppo_experience_kernel compilation failed (is NVRTC available?): {e}"
))
})?;
let module = context.load_module(ptx).map_err(|e| {
MLError::ModelError(format!("Failed to load ppo_experience_kernel module: {e}"))
})?;

View File

@@ -85,6 +85,56 @@ pub fn optimal_launch_dims(n_items: u32, max_threads_per_block: u32) -> (u32, u3
(grid, block)
}
/// Compile CUDA source to PTX with architecture-aware options.
///
/// When a CUDA device is available, queries its compute capability and passes
/// `-arch=compute_XX` to NVRTC. On Hopper (sm_90+), this enables TMA
/// (`cp.async.bulk`) tile loads via `__CUDA_ARCH__ >= 900` guards in the
/// kernel source. Falls back to `compile_ptx` (default arch) when the
/// device capability cannot be determined.
///
/// # Feature gate
/// Only available with the `cuda` feature.
#[cfg(feature = "cuda")]
pub fn compile_ptx_for_device(
src: &str,
context: &candle_core::cuda_backend::cudarc::driver::CudaContext,
) -> Result<candle_core::cuda_backend::cudarc::nvrtc::Ptx, String> {
use candle_core::cuda_backend::cudarc;
use cudarc::driver::sys::CUdevice_attribute;
let major = context
.attribute(CUdevice_attribute::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR)
.unwrap_or(0);
let minor = context
.attribute(CUdevice_attribute::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR)
.unwrap_or(0);
// Build arch string: "compute_90", "compute_80", etc.
// NVRTC accepts the compute_XX form (virtual architecture) which enables
// __CUDA_ARCH__ and arch-specific code generation.
let arch_str: &'static str = match (major, minor) {
(9, 0) => "compute_90",
(9, _) => "compute_90", // sm_90a and beyond
(8, 9) => "compute_89",
(8, 6) => "compute_86",
(8, 0) => "compute_80",
(7, 5) => "compute_75",
(7, 0) => "compute_70",
_ => {
// Unknown or very old/new — use default NVRTC behavior.
return cudarc::nvrtc::compile_ptx(src)
.map_err(|e| e.to_string());
}
};
let opts = cudarc::nvrtc::CompileOptions {
arch: Some(arch_str),
..Default::default()
};
cudarc::nvrtc::compile_ptx_with_opts(src, opts)
.map_err(|e| e.to_string())
}
/// Pre-uploaded GPU training data for DQN trainer.
///
/// Holds market features [N, 42] and target prices [N, 4] as GPU tensors.