| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| #![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")); |
|
|
| |
| #[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 {} |
|
|
| |
| #[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<u32> { |
| std::env::var("HTM_FUSED_GRID_CAP") |
| .ok() |
| .and_then(|s| s.parse::<u32>().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<u32>, |
| ) -> 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); |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| pub struct FusedState { |
| dev: Arc<CudaDevice>, |
| raw_kernel: RawFusedKernel, |
|
|
| pub inhibition_threshold: CudaSlice<f32>, |
| pub cell_active_bits_a: CudaSlice<u32>, |
| pub cell_active_bits_b: CudaSlice<u32>, |
| pub cell_winner_bits_a: CudaSlice<u32>, |
| pub cell_winner_bits_b: CudaSlice<u32>, |
| pub barrier_counters: CudaSlice<u32>, |
| pub step_scratch: CudaSlice<u32>, |
|
|
| pub grid_dim_x: u32, |
| pub block_dim_x: u32, |
| pub cooperative_grid_limit: u32, |
| pub launch_mode: FusedLaunchMode, |
| pub iter_counter: u32, |
|
|
| |
| #[allow(dead_code)] |
| pub initial_threshold: f32, |
| } |
|
|
| impl FusedState { |
| pub fn new( |
| dev: Arc<CudaDevice>, |
| n_columns: usize, |
| cells_per_column: usize, |
| initial_threshold: f32, |
| ) -> Result<Self, DriverError> { |
| 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::<f32>(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::<u32>(bits_words)?; |
| let cell_active_bits_b = dev.alloc_zeros::<u32>(bits_words)?; |
| let cell_winner_bits_a = dev.alloc_zeros::<u32>(bits_words)?; |
| let cell_winner_bits_b = dev.alloc_zeros::<u32>(bits_words)?; |
| let barrier_counters = dev.alloc_zeros::<u32>(3)?; |
| let step_scratch = dev.alloc_zeros::<u32>(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()) |
| }?; |
|
|
| |
| 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, |
| }) |
| } |
|
|
| |
| 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)?; |
| |
| |
| Ok(()) |
| } |
| } |
|
|
| |
| #[allow(clippy::too_many_arguments)] |
| pub fn launch_fused( |
| sp: &mut SpatialPoolerGpu, |
| tm: &mut TemporalMemoryGpu, |
| fused: &mut FusedState, |
| inputs_flat: &CudaSlice<u8>, |
| cols_out: &mut CudaSlice<u8>, |
| anom_out: &mut CudaSlice<f32>, |
| t: usize, |
| input_bits: usize, |
| learn: bool, |
| ) -> Result<(), DriverError> { |
| |
| 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(()) |
| } |
|
|