//! Fused HTM megakernel launcher. //! //! Collapses the 12-kernel per-timestep pipeline (and the outer T-loop) into //! a single kernel launch per forward. See `kernels/htm_fused_step.cu` for //! the kernel design and the cross-block coherence strategy (grid barrier //! via device counter with all blocks concurrently resident). //! //! Launch invariant: `grid_dim.x <= concurrent-block capacity`. Host code //! probes the device SM count at construction and caps grid_dim.x //! accordingly — otherwise the grid barrier deadlocks. //! //! Semantic change from the top-K pipeline: activation is per-column //! threshold-based (local lateral inhibition) instead of global top-K. //! A per-column `inhibition_threshold` is tracked and EMA-steered to hit //! the sparsity target. This is a real architectural change and is //! documented in `docs/GPU_HTM.md`. #![cfg(feature = "gpu")] use std::ffi::CString; use std::sync::Arc; use cudarc::driver::{result, sys, CudaDevice, CudaSlice, DeviceRepr, DevicePtr, DriverError, LaunchConfig}; use cudarc::nvrtc::Ptx; use super::sp_gpu::SpatialPoolerGpu; use super::tm_gpu::{TemporalMemoryGpu, MAX_SEGMENTS_PER_CELL, MAX_SYN_PER_SEGMENT}; const PTX_HTM_FUSED: &str = include_str!(concat!(env!("HTM_GPU_PTX_DIR"), "/htm_fused_step.ptx")); /// Struct-by-value pointer pack — matches C-side `FusedPtrs`. #[repr(C)] #[derive(Clone, Copy)] pub struct FusedPtrs { pub syn_bit: u64, pub syn_perm: u64, pub boost: u64, pub active_duty: u64, pub inhibition_threshold: u64, pub seg_cell_id: u64, pub seg_syn_count: u64, pub syn_presyn: u64, pub tm_syn_perm: u64, pub cell_seg_count: u64, pub cell_active_a: u64, pub cell_active_b: u64, pub cell_winner_a: u64, pub cell_winner_b: u64, pub inputs: u64, pub cols_out: u64, pub anom_out: u64, pub barrier_counters: u64, pub step_scratch: u64, } unsafe impl DeviceRepr for FusedPtrs {} /// Launch-time config — matches C-side `FusedConfig` 1:1. #[repr(C)] #[derive(Clone, Copy)] pub struct FusedConfig { pub input_bits: u32, pub n_columns: u32, pub synapses_per_col: u32, pub conn_thr: f32, pub sp_inc: f32, pub sp_dec: f32, pub sparsity_target: f32, pub duty_alpha: f32, pub thr_adapt_rate: f32, pub cells_per_column: u32, pub n_cells: u32, pub bits_words: u32, pub max_segments_per_cell: u32, pub synapses_per_segment: u32, pub activation_threshold: u32, pub learning_threshold: u32, pub max_new_synapses: u32, pub conn_thr_i16: i32, pub perm_inc_i16: i32, pub perm_dec_i16: i32, pub predicted_seg_dec_i16: i32, pub initial_perm_i16: i32, pub t: u32, pub learn: u32, pub iter_seed: u32, pub cooperative_grid_sync: u32, } unsafe impl DeviceRepr for FusedConfig {} #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum FusedLaunchMode { Cooperative, SoftwareBarrier, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) struct FusedLaunchPlan { pub mode: FusedLaunchMode, pub grid_dim_x: u32, pub block_dim_x: u32, pub cooperative_grid_limit: u32, pub sm_count: u32, } impl FusedLaunchPlan { pub(crate) fn uses_cooperative_launch(&self) -> bool { matches!(self.mode, FusedLaunchMode::Cooperative) } } fn fused_grid_cap_override() -> Option { std::env::var("HTM_FUSED_GRID_CAP") .ok() .and_then(|s| s.parse::().ok()) .map(|v| v.max(1)) } pub(crate) fn plan_fused_launch( sm_count: u32, cooperative_supported: bool, cooperative_grid_limit: u32, grid_cap_override: Option, ) -> FusedLaunchPlan { let sm_count = sm_count.max(1); let block_dim_x = 1024u32; if cooperative_supported && cooperative_grid_limit > 0 { let default_grid_cap = 16u32; let grid_cap = grid_cap_override.unwrap_or(default_grid_cap); return FusedLaunchPlan { mode: FusedLaunchMode::Cooperative, grid_dim_x: cooperative_grid_limit.min(grid_cap).max(1), block_dim_x, cooperative_grid_limit, sm_count, }; } let default_grid_cap = 8u32; let grid_cap = grid_cap_override.unwrap_or(default_grid_cap); FusedLaunchPlan { mode: FusedLaunchMode::SoftwareBarrier, grid_dim_x: sm_count.min(grid_cap).max(1), block_dim_x, cooperative_grid_limit: 0, sm_count, } } struct RawFusedKernel { module: sys::CUmodule, function: sys::CUfunction, } unsafe impl Send for RawFusedKernel {} unsafe impl Sync for RawFusedKernel {} impl Drop for RawFusedKernel { fn drop(&mut self) { unsafe { let _ = result::module::unload(self.module); } } } /// Owns fused-path-only device state: /// - per-column inhibition threshold (replaces global top-K) /// - ping-pong cell_active/cell_winner bitsets /// - 3-slot rotating grid barrier counters /// - step_scratch (n_active, n_unpred per timestep) pub struct FusedState { dev: Arc, raw_kernel: RawFusedKernel, pub inhibition_threshold: CudaSlice, pub cell_active_bits_a: CudaSlice, pub cell_active_bits_b: CudaSlice, pub cell_winner_bits_a: CudaSlice, pub cell_winner_bits_b: CudaSlice, pub barrier_counters: CudaSlice, // length 3 pub step_scratch: CudaSlice, // length 6 pub grid_dim_x: u32, pub block_dim_x: u32, pub cooperative_grid_limit: u32, pub launch_mode: FusedLaunchMode, pub iter_counter: u32, // Config mirror (read-only after init). #[allow(dead_code)] pub initial_threshold: f32, } impl FusedState { pub fn new( dev: Arc, n_columns: usize, cells_per_column: usize, initial_threshold: f32, ) -> Result { let n_cells = n_columns * cells_per_column; assert!(n_cells % 32 == 0, "n_cells must be divisible by 32 for bitsets"); let bits_words = n_cells / 32; let mut inhibition_threshold = dev.alloc_zeros::(n_columns)?; let init_vec = vec![initial_threshold; n_columns]; dev.htod_sync_copy_into(&init_vec, &mut inhibition_threshold)?; let cell_active_bits_a = dev.alloc_zeros::(bits_words)?; let cell_active_bits_b = dev.alloc_zeros::(bits_words)?; let cell_winner_bits_a = dev.alloc_zeros::(bits_words)?; let cell_winner_bits_b = dev.alloc_zeros::(bits_words)?; let barrier_counters = dev.alloc_zeros::(3)?; let step_scratch = dev.alloc_zeros::(6)?; unsafe { result::ctx::set_current(*dev.cu_primary_ctx())?; } if dev.get_func("htm_fused", "htm_fused_step").is_none() { dev.load_ptx(Ptx::from_src(PTX_HTM_FUSED), "htm_fused", &["htm_fused_step"])?; } let ptx = CString::new(PTX_HTM_FUSED).expect("PTX contains no interior nul bytes"); let module = unsafe { result::module::load_data(ptx.as_ptr().cast()) }?; let function = unsafe { result::module::get_function(module, CString::new("htm_fused_step").unwrap()) }?; // Probe SM count. let sm_count = match dev.attribute( cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT, ) { Ok(v) => v as u32, Err(_) => 16u32, }; let cooperative_supported = matches!( dev.attribute(sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_COOPERATIVE_LAUNCH), Ok(v) if v > 0 ); let cooperative_grid_limit = if cooperative_supported { let blocks_per_sm = unsafe { result::occupancy::max_active_block_per_multiprocessor(function, 1024, 0) } .ok() .map(|v| v.max(0) as u32) .unwrap_or(0); sm_count.saturating_mul(blocks_per_sm) } else { 0 }; let launch_plan = plan_fused_launch( sm_count, cooperative_supported, cooperative_grid_limit, fused_grid_cap_override(), ); Ok(Self { dev, raw_kernel: RawFusedKernel { module, function }, inhibition_threshold, cell_active_bits_a, cell_active_bits_b, cell_winner_bits_a, cell_winner_bits_b, barrier_counters, step_scratch, grid_dim_x: launch_plan.grid_dim_x, block_dim_x: launch_plan.block_dim_x, cooperative_grid_limit: launch_plan.cooperative_grid_limit, launch_mode: launch_plan.mode, iter_counter: 0, initial_threshold, }) } /// Reset fused state. Called at region.reset(). pub fn reset(&mut self) -> Result<(), DriverError> { self.dev.memset_zeros(&mut self.cell_active_bits_a)?; self.dev.memset_zeros(&mut self.cell_active_bits_b)?; self.dev.memset_zeros(&mut self.cell_winner_bits_a)?; self.dev.memset_zeros(&mut self.cell_winner_bits_b)?; self.dev.memset_zeros(&mut self.barrier_counters)?; self.dev.memset_zeros(&mut self.step_scratch)?; // Do NOT reset inhibition_threshold — it's learned state. A hard // reset of TM state should NOT forget the sparsity calibration. Ok(()) } } /// Launch the fused megakernel. Processes all T timesteps in one kernel. #[allow(clippy::too_many_arguments)] pub fn launch_fused( sp: &mut SpatialPoolerGpu, tm: &mut TemporalMemoryGpu, fused: &mut FusedState, inputs_flat: &CudaSlice, cols_out: &mut CudaSlice, anom_out: &mut CudaSlice, t: usize, input_bits: usize, learn: bool, ) -> Result<(), DriverError> { // Reset barrier counters + scratch before each launch (safe re-entry). sp.dev_ref().memset_zeros(&mut fused.barrier_counters)?; sp.dev_ref().memset_zeros(&mut fused.step_scratch)?; fused.iter_counter = fused.iter_counter.wrapping_add(1); let cfg = FusedConfig { input_bits: input_bits as u32, n_columns: sp.n_columns_accessor() as u32, synapses_per_col: sp.synapses_per_col_accessor() as u32, conn_thr: sp.conn_thr_accessor(), sp_inc: sp.inc_accessor(), sp_dec: sp.dec_accessor(), sparsity_target: sp.sparsity_accessor(), duty_alpha: 1.0f32 / sp.duty_period_accessor().max(1.0), thr_adapt_rate: 0.001f32, cells_per_column: tm.cells_per_column as u32, n_cells: tm.n_cells as u32, bits_words: tm.bits_words as u32, max_segments_per_cell: MAX_SEGMENTS_PER_CELL as u32, synapses_per_segment: MAX_SYN_PER_SEGMENT as u32, activation_threshold: tm.activation_threshold, learning_threshold: tm.learning_threshold, max_new_synapses: tm.max_new_synapse_count, conn_thr_i16: tm.conn_thr_i16 as i32, perm_inc_i16: tm.perm_inc_i16 as i32, perm_dec_i16: tm.perm_dec_i16 as i32, predicted_seg_dec_i16: tm.predicted_seg_dec_i16 as i32, initial_perm_i16: tm.initial_perm_i16 as i32, t: t as u32, learn: if learn { 1 } else { 0 }, iter_seed: fused.iter_counter, cooperative_grid_sync: if fused.launch_mode == FusedLaunchMode::Cooperative { 1 } else { 0 }, }; let ptrs = FusedPtrs { syn_bit: *sp.syn_bit_accessor().device_ptr(), syn_perm: *sp.syn_perm_accessor().device_ptr(), boost: *sp.boost_accessor().device_ptr(), active_duty: *sp.active_duty_accessor().device_ptr(), inhibition_threshold: *fused.inhibition_threshold.device_ptr(), seg_cell_id: *tm.seg_cell_id_accessor().device_ptr(), seg_syn_count: *tm.seg_syn_count_accessor().device_ptr(), syn_presyn: *tm.syn_presyn_accessor().device_ptr(), tm_syn_perm: *tm.syn_perm_accessor().device_ptr(), cell_seg_count: *tm.cell_seg_count_accessor().device_ptr(), cell_active_a: *fused.cell_active_bits_a.device_ptr(), cell_active_b: *fused.cell_active_bits_b.device_ptr(), cell_winner_a: *fused.cell_winner_bits_a.device_ptr(), cell_winner_b: *fused.cell_winner_bits_b.device_ptr(), inputs: *inputs_flat.device_ptr(), cols_out: *cols_out.device_ptr(), anom_out: *anom_out.device_ptr(), barrier_counters: *fused.barrier_counters.device_ptr(), step_scratch: *fused.step_scratch.device_ptr(), }; let launch_cfg = LaunchConfig { grid_dim: (fused.grid_dim_x, 1, 1), block_dim: (fused.block_dim_x, 1, 1), shared_mem_bytes: 0, }; unsafe { result::ctx::set_current(*sp.dev_ref().cu_primary_ctx())?; let mut kernel_params: [*mut std::ffi::c_void; 2] = [ (&ptrs as *const FusedPtrs).cast_mut().cast(), (&cfg as *const FusedConfig).cast_mut().cast(), ]; match fused.launch_mode { FusedLaunchMode::Cooperative => result::launch_cooperative_kernel( fused.raw_kernel.function, launch_cfg.grid_dim, launch_cfg.block_dim, launch_cfg.shared_mem_bytes, *sp.dev_ref().cu_stream(), &mut kernel_params, )?, FusedLaunchMode::SoftwareBarrier => result::launch_kernel( fused.raw_kernel.function, launch_cfg.grid_dim, launch_cfg.block_dim, launch_cfg.shared_mem_bytes, *sp.dev_ref().cu_stream(), &mut kernel_params, )?, } } Ok(()) }