From 542be32bedf4233f56dff8d302f3ac5f4f57de09 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 8 Mar 2026 02:10:38 +0100 Subject: [PATCH] perf(ml): pad DQN input dims for H100 tensor core alignment MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit cuBLAS requires M/N/K dimensions to be multiples of 8 for BF16 tensor core dispatch. state_dim=51 (with OFI) caused silent fallback to scalar FMA ops, leaving tensor cores at 0.1% utilization despite BF16 enabled. Pad state_dim to next 8-multiple (51→56, 43→48) at network construction and zero-pad input tensors in forward pass across all DQN network types: Sequential, DuelingQNetwork, DistributionalDuelingQNetwork, NetworkLayers. Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/dqn/distributional_dueling.rs | 14 +++++++-- crates/ml/src/dqn/dqn.rs | 22 +++++++++++-- crates/ml/src/dqn/dueling.rs | 15 +++++++-- crates/ml/src/dqn/mixed_precision.rs | 35 ++++++++++++++++++++- crates/ml/src/dqn/network.rs | 14 ++++++--- 5 files changed, 87 insertions(+), 13 deletions(-) diff --git a/crates/ml/src/dqn/distributional_dueling.rs b/crates/ml/src/dqn/distributional_dueling.rs index 8f770f83d..999883c13 100644 --- a/crates/ml/src/dqn/distributional_dueling.rs +++ b/crates/ml/src/dqn/distributional_dueling.rs @@ -148,6 +148,8 @@ pub struct DistributionalDuelingQNetwork { /// Device (CPU or CUDA) device: Device, + /// Padded input dim for tensor core alignment + padded_state_dim: usize, } impl DistributionalDuelingQNetwork { @@ -162,13 +164,17 @@ impl DistributionalDuelingQNetwork { /// /// New DistributionalDuelingQNetwork instance with Xavier-initialized weights pub fn new(config: DistributionalDuelingConfig, device: Device) -> Result { + use super::mixed_precision::align_dim_for_tensor_cores; + + let padded_state_dim = align_dim_for_tensor_cores(config.state_dim, &device); + let vars = VarMap::new(); let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); // Build shared feature layers with RMSNorm after each let mut shared_layers = Vec::new(); let mut shared_norms = Vec::new(); - let mut current_dim = config.state_dim; + let mut current_dim = padded_state_dim; for (i, &hidden_dim) in config.shared_hidden_dims.iter().enumerate() { let layer_name = format!("shared_{}", i); @@ -229,6 +235,7 @@ impl DistributionalDuelingQNetwork { config, vars, device, + padded_state_dim, }) } @@ -259,8 +266,9 @@ impl DistributionalDuelingQNetwork { .dim(0) .map_err(|e| MLError::ModelError(format!("Failed to get batch size: {}", e)))?; - // Shared feature extraction: Linear → LeakyReLU → RMSNorm - let mut h = state.clone(); + // Zero-pad input for tensor core alignment, then shared feature extraction + let mut h = super::mixed_precision::pad_to_aligned(&state, self.padded_state_dim) + .map_err(|e| MLError::ModelError(format!("Tensor core padding failed: {}", e)))?; for (i, (layer, norm)) in self .shared_layers .iter() diff --git a/crates/ml/src/dqn/dqn.rs b/crates/ml/src/dqn/dqn.rs index db30b7608..a2d7c90ac 100644 --- a/crates/ml/src/dqn/dqn.rs +++ b/crates/ml/src/dqn/dqn.rs @@ -851,6 +851,9 @@ pub struct Sequential { device: Device, vars: VarMap, leaky_relu_alpha: f64, + /// Padded input dimension (aligned to 8 for tensor core matmul). + /// When > raw input_dim, forward() zero-pads the input tensor. + padded_input_dim: usize, } impl Sequential { @@ -864,12 +867,23 @@ impl Sequential { use_noisy_nets: bool, noisy_sigma_init: f64, ) -> Result { + use super::mixed_precision::align_dim_for_tensor_cores; + + let padded_input_dim = align_dim_for_tensor_cores(input_dim, &device); + if padded_input_dim != input_dim { + tracing::info!( + "Tensor core alignment: input_dim={} → padded={} (K dim aligned to 8)", + input_dim, + padded_input_dim + ); + } + let vars = VarMap::new(); let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); let mut layers = Vec::new(); let mut noisy_layers = Vec::new(); - let mut current_dim = input_dim; + let mut current_dim = padded_input_dim; if use_noisy_nets { // Create NoisyLinear layers (Rainbow DQN exploration) @@ -916,6 +930,7 @@ impl Sequential { device, vars, leaky_relu_alpha, + padded_input_dim, }) } @@ -923,7 +938,10 @@ impl Sequential { pub fn forward(&self, input: &Tensor) -> Result { let input = crate::dqn::mixed_precision::ensure_training_dtype(input) .map_err(|e| MLError::ModelError(e.to_string()))?; - let mut x = input.clone(); + + // Zero-pad input to aligned dimension for tensor core utilization + let mut x = super::mixed_precision::pad_to_aligned(&input, self.padded_input_dim) + .map_err(|e| MLError::ModelError(format!("Tensor core padding failed: {}", e)))?; if self.use_noisy_nets { // Use noisy layers (Rainbow DQN exploration) diff --git a/crates/ml/src/dqn/dueling.rs b/crates/ml/src/dqn/dueling.rs index e2bfb4634..b29db576d 100644 --- a/crates/ml/src/dqn/dueling.rs +++ b/crates/ml/src/dqn/dueling.rs @@ -122,6 +122,9 @@ pub struct DuelingQNetwork { /// Device (CPU or CUDA) device: Device, + + /// Padded input dim for tensor core alignment + padded_state_dim: usize, } impl DuelingQNetwork { @@ -136,12 +139,16 @@ impl DuelingQNetwork { /// /// New DuelingQNetwork instance with Xavier-initialized weights pub fn new(config: DuelingConfig, device: Device) -> Result { + use super::mixed_precision::align_dim_for_tensor_cores; + + let padded_state_dim = align_dim_for_tensor_cores(config.state_dim, &device); + let vars = VarMap::new(); let var_builder = VarBuilder::from_varmap(&vars, training_dtype(&device), &device); // Build shared feature layers let mut shared_layers = Vec::new(); - let mut current_dim = config.state_dim; + let mut current_dim = padded_state_dim; for (i, &hidden_dim) in config.shared_hidden_dims.iter().enumerate() { let layer_name = format!("shared_{}", i); @@ -184,6 +191,7 @@ impl DuelingQNetwork { config, vars, device, + padded_state_dim, }) } @@ -208,8 +216,9 @@ impl DuelingQNetwork { pub fn forward(&self, state: &Tensor) -> Result { let state = crate::dqn::mixed_precision::ensure_training_dtype(state) .map_err(|e| MLError::ModelError(e.to_string()))?; - // Shared feature extraction - let mut h = state.clone(); + // Zero-pad input for tensor core alignment + let mut h = super::mixed_precision::pad_to_aligned(&state, self.padded_state_dim) + .map_err(|e| MLError::ModelError(format!("Tensor core padding failed: {}", e)))?; for (i, layer) in self.shared_layers.iter().enumerate() { h = layer.forward(&h).map_err(|e| { MLError::ModelError(format!("Shared layer {} forward failed: {}", i, e)) diff --git a/crates/ml/src/dqn/mixed_precision.rs b/crates/ml/src/dqn/mixed_precision.rs index 676940178..4552d8fd8 100644 --- a/crates/ml/src/dqn/mixed_precision.rs +++ b/crates/ml/src/dqn/mixed_precision.rs @@ -7,11 +7,44 @@ //! //! Wave 26 P2.1: AMP implementation for production DQN training -use candle_core::{DType, Tensor}; +use candle_core::{DType, Device, Tensor}; use serde::{Deserialize, Serialize}; use crate::MLError; +/// Round up a dimension to the next multiple of 8 for tensor core alignment. +/// +/// BF16/FP16 tensor cores (Volta+) require M, N, K dimensions to be multiples +/// of 8 for HMMA (half-precision matrix multiply-accumulate) instructions. +/// Unaligned dimensions cause cuBLAS to fall back to scalar FMA ops, leaving +/// tensor cores completely idle (0% utilization). +/// +/// Only pads on CUDA devices where tensor cores are available. +/// Returns the original dimension on CPU/Metal. +#[inline] +pub fn align_dim_for_tensor_cores(dim: usize, device: &Device) -> usize { + match device { + Device::Cuda(_) => (dim + 7) & !7, + _ => dim, + } +} + +/// Zero-pad a 2D tensor's last dimension to `target_dim` if needed. +/// +/// Used at network entry points to align the K dimension for tensor core matmul. +/// When `input.dim(-1) == target_dim`, returns the input unchanged (no allocation). +pub fn pad_to_aligned(input: &Tensor, target_dim: usize) -> Result { + let current_dim = input.dim(candle_core::D::Minus1)?; + if current_dim >= target_dim { + return Ok(input.clone()); + } + let batch_dims: Vec = input.dims()[..input.dims().len() - 1].to_vec(); + let mut pad_shape = batch_dims; + pad_shape.push(target_dim - current_dim); + let padding = Tensor::zeros(&pad_shape[..], input.dtype(), input.device())?; + Tensor::cat(&[input, &padding], input.dims().len() - 1) +} + /// Configuration for mixed precision training #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MixedPrecisionConfig { diff --git a/crates/ml/src/dqn/network.rs b/crates/ml/src/dqn/network.rs index 8eded805d..43a8979ff 100644 --- a/crates/ml/src/dqn/network.rs +++ b/crates/ml/src/dqn/network.rs @@ -145,6 +145,7 @@ struct NetworkLayers { layers: Vec, dropout: Dropout, training: bool, + padded_input_dim: usize, } impl NetworkLayers { @@ -164,8 +165,12 @@ impl NetworkLayers { dropout_rate: f64, training: bool, ) -> CandleResult { + let padded_input_dim = crate::dqn::mixed_precision::align_dim_for_tensor_cores( + config.state_dim, _device, + ); + let mut layers = Vec::new(); - let mut input_dim = config.state_dim; + let mut input_dim = padded_input_dim; // Create hidden layers with Xavier initialization for (i, &hidden_dim) in config.hidden_dims.iter().enumerate() { @@ -185,7 +190,7 @@ impl NetworkLayers { let dropout = Dropout::new(dropout_rate as f32); - Ok(Self { layers, dropout, training }) + Ok(Self { layers, dropout, training, padded_input_dim }) } /// Forward pass with optional mixed precision (BF16/FP16). @@ -207,7 +212,8 @@ impl NetworkLayers { _ => (xs.clone(), false), }; - let mut x = x_input; + // Zero-pad input for tensor core alignment + let mut x = crate::dqn::mixed_precision::pad_to_aligned(&x_input, self.padded_input_dim)?; for (i, layer) in self.layers.iter().enumerate() { x = layer.forward(&x)?; @@ -224,7 +230,7 @@ impl NetworkLayers { impl Module for NetworkLayers { fn forward(&self, xs: &Tensor) -> CandleResult { let xs = crate::dqn::mixed_precision::ensure_training_dtype(xs)?; - let mut x = xs.clone(); + let mut x = crate::dqn::mixed_precision::pad_to_aligned(&xs, self.padded_input_dim)?; // Forward through hidden layers with LeakyReLU activation and dropout // LeakyReLU prevents dead neurons (0.01 gradient for negative inputs vs 0 for ReLU)