1use crate::{
2 WmmaCompiler,
3 compute::{CudaServer, context::CudaContext},
4 device::CudaDevice,
5};
6use cubecl_common::{
7 device::{Device, DeviceService},
8 profile::TimingMethod,
9};
10use cubecl_core::{
11 MemoryConfiguration, Runtime,
12 device::{DeviceId, ServerUtilitiesHandle},
13 ir::{
14 BarrierLevel, ContiguousElements, DeviceIdentity, DeviceProperties, ElemType, FloatKind,
15 HardwareProperties, MatrixLayout, MemoryDeviceProperties, MmaProperties, OpaqueType,
16 StorageType, TargetProperties, Type, VectorSize,
17 features::{AtomicUsage, Plane, Tma, TypeUsage},
18 },
19 server::ServerUtilities,
20 zspace::{Shape, Strides, striding::has_pitched_row_major_strides},
21};
22use cubecl_cpp::{
23 ComputeKernel, DialectWmmaCompiler,
24 cuda::{CudaDialect, arch::CudaArchitecture, mma::contiguous_elements_cuda},
25 register_supported_types,
26 shared::{
27 CompilationOptions, CppCompiler, CppSupportedFeatures, register_mma_features,
28 register_scaled_mma_features, register_wmma_features,
29 },
30};
31use cubecl_runtime::{
32 allocator::PitchedMemoryLayoutPolicy, client::ComputeClient, logging::ServerLogger,
33};
34use cudarc::driver::sys::{CUDA_VERSION, cuDeviceTotalMem_v2};
35use std::{mem::MaybeUninit, sync::Arc};
36
37#[derive(Default)]
39pub struct RuntimeOptions {
40 pub memory_config: MemoryConfiguration,
42}
43
44#[derive(Debug, Clone)]
45pub struct CudaRuntime;
46
47impl DeviceService for CudaServer {
48 fn init(device_id: cubecl_common::device::DeviceId) -> Self {
49 let options = RuntimeOptions::default();
50 let device = CudaDevice::from_id(device_id);
51
52 cudarc::driver::result::init().unwrap();
54 let device_index = device.index as i32;
55 let device_ptr = cudarc::driver::result::device::get(device_index).unwrap();
56 let arch_major;
57 let arch_version = unsafe {
60 arch_major = cudarc::driver::result::device::get_attribute(
61 device_ptr,
62 cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
63 )
64 .unwrap();
65 let minor = cudarc::driver::result::device::get_attribute(
66 device_ptr,
67 cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
68 )
69 .unwrap();
70 arch_major * 10 + minor
71 } as u32;
72
73 let mem_alignment = 512;
77
78 let arch = CudaArchitecture {
80 version: arch_version,
81 };
82 let supported_wmma_combinations = WmmaCompiler::supported_wmma_combinations(&arch);
83 let supported_mma_combinations = WmmaCompiler::supported_mma_combinations(&arch);
84 let supported_scaled_mma_combinations =
85 WmmaCompiler::supported_scaled_mma_combinations(&arch);
86
87 let ctx = unsafe {
90 let ctx = cudarc::driver::result::primary_ctx::retain(device_ptr).unwrap();
91 cudarc::driver::result::ctx::set_current(ctx).unwrap();
92 ctx
93 };
94
95 let max_memory = unsafe {
98 let mut bytes = MaybeUninit::uninit();
99 cuDeviceTotalMem_v2(bytes.as_mut_ptr(), device_ptr);
100 bytes.assume_init() as u64
101 };
102 let mem_properties = MemoryDeviceProperties {
103 max_page_size: max_memory / 4,
104 alignment: mem_alignment as u64,
105 };
106
107 let mut comp_opts = CompilationOptions {
108 supports_features: CppSupportedFeatures {
109 fast_math: true,
110 ..Default::default()
111 },
112 ..Default::default()
113 };
114
115 let hardware_props = unsafe {
118 use cudarc::driver::{result::device::get_attribute, sys::CUdevice_attribute::*};
119 let warp_size =
120 get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_WARP_SIZE).unwrap() as u32;
121 let max_shared = get_attribute(
122 device_ptr,
123 CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK_OPTIN,
124 )
125 .unwrap() as usize;
126 let max_threads = get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_THREADS_PER_BLOCK)
127 .unwrap() as u32;
128 let block_dim_x =
129 get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_X).unwrap();
130 let block_dim_y =
131 get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Y).unwrap();
132 let block_dim_z =
133 get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_BLOCK_DIM_Z).unwrap();
134 let max_cube_dim = (block_dim_x as u32, block_dim_y as u32, block_dim_z as u32);
135
136 let grid_dim_x = get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_X).unwrap();
137 let grid_dim_y = get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Y).unwrap();
138 let grid_dim_z = get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_Z).unwrap();
139 let max_cube_count = (grid_dim_x as u32, grid_dim_y as u32, grid_dim_z as u32);
140
141 let num_streaming_multiprocessors = Some(
142 get_attribute(device_ptr, CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap() as u32,
143 );
144 let num_tensor_cores = tensor_cores_per_sm(arch_version);
145
146 comp_opts.warp_size = warp_size;
147
148 HardwareProperties {
149 load_width: 128,
150 plane_size_min: warp_size,
151 plane_size_max: warp_size,
152 max_bindings: crate::device::CUDA_MAX_BINDINGS,
153 max_shared_memory_size: max_shared,
154 max_cube_count,
155 max_units_per_cube: max_threads,
156 max_cube_dim,
157 num_streaming_multiprocessors,
158 num_tensor_cores,
159 min_tensor_cores_dim: if supported_wmma_combinations.is_empty() {
160 None
161 } else {
162 Some(8)
163 },
164 num_cpu_cores: None,
165 max_vector_size: VectorSize::MAX,
166 cube_mma_reserved_shared_memory: 0,
167 }
168 };
169
170 let fingerprint = format!("ptx_sm{arch_version}");
174 let device_name = cudarc::driver::result::device::get_name(device_ptr)
177 .unwrap_or_else(|_| "unknown CUDA device".to_string());
178
179 let mut device_props = DeviceProperties::new(
180 Default::default(),
181 mem_properties.clone(),
182 hardware_props,
183 TimingMethod::System,
184 DeviceIdentity {
185 name: device_name,
186 fingerprint: fingerprint.clone(),
187 },
188 );
189 register_supported_types(&mut device_props);
190 device_props.register_type_usage(ElemType::Float(FloatKind::TF32), TypeUsage::Conversion);
191 if arch_version >= 60 {
192 device_props.register_atomic_type_usage(
193 Type::atomic(ElemType::Float(FloatKind::F64)),
194 AtomicUsage::Add | AtomicUsage::LoadStore,
195 );
196 }
197 if arch_version >= 70 {
198 device_props.register_atomic_type_usage(
199 Type::atomic(ElemType::Float(FloatKind::F16)),
200 AtomicUsage::Add,
201 );
202 device_props.register_atomic_type_usage(
203 Type::atomic(Type::scalar(ElemType::Float(FloatKind::F16)).with_vector_size(2)),
204 AtomicUsage::Add | AtomicUsage::LoadStore,
205 );
206 device_props.register_opaque_type(OpaqueType::Barrier(BarrierLevel::Unit));
207 device_props.register_opaque_type(OpaqueType::Barrier(BarrierLevel::Cube));
208 device_props.features.plane.insert(Plane::Sync);
209 comp_opts.supports_features.grid_constants = true;
210 }
211
212 if arch_version >= 75 {
213 device_props
214 .features
215 .matmul
216 .ldmatrix
217 .insert(ElemType::Float(FloatKind::F16).into());
218 device_props
219 .features
220 .matmul
221 .ldmatrix
222 .insert(ElemType::Float(FloatKind::BF16).into());
223 comp_opts.supports_features.fast_tanh = CUDA_VERSION >= 12080;
224 }
225
226 if arch_version >= 80 {
227 device_props.features.copy_async = true;
228 }
229
230 if arch_version >= 89 {
236 device_props.register_type_usage(
237 ElemType::Float(FloatKind::E4M3),
238 TypeUsage::Conversion | TypeUsage::Buffer,
239 );
240 device_props.register_type_usage(
241 ElemType::Float(FloatKind::E5M2),
242 TypeUsage::Conversion | TypeUsage::Buffer,
243 );
244 }
245 if arch_version >= 90 {
246 device_props.features.tma.insert(Tma::Base);
247 device_props.register_opaque_type(OpaqueType::TensorMap);
248 device_props.features.cube_cluster = true;
249 comp_opts.supports_features.clusters = true;
250 comp_opts.supports_features.elect_sync = true;
251 device_props
252 .features
253 .matmul
254 .stmatrix
255 .insert(ElemType::Float(FloatKind::F16).into());
256 device_props
257 .features
258 .matmul
259 .stmatrix
260 .insert(ElemType::Float(FloatKind::BF16).into());
261
262 if CUDA_VERSION > 12080 {
263 device_props.register_atomic_type_usage(
264 Type::atomic(Type::scalar(ElemType::Float(FloatKind::F32)).with_vector_size(2)),
265 AtomicUsage::LoadStore | AtomicUsage::Add,
266 );
267 device_props.register_atomic_type_usage(
268 Type::atomic(Type::scalar(ElemType::Float(FloatKind::F32)).with_vector_size(4)),
269 AtomicUsage::LoadStore | AtomicUsage::Add,
270 );
271 }
272 }
273
274 if arch_version >= 100 {
275 device_props.features.tma.insert(Tma::Im2colWide);
276 }
281
282 if arch_major == 10 || arch_major == 11 || arch_major == 12 {
286 device_props
287 .register_type_usage(ElemType::Float(FloatKind::E2M1), TypeUsage::Conversion);
288 device_props.register_type_usage(
289 StorageType::Packed(ElemType::Float(FloatKind::E2M1), 2),
290 TypeUsage::Conversion | TypeUsage::Buffer,
291 );
292 device_props.register_type_usage(
293 ElemType::Float(FloatKind::E2M3),
294 TypeUsage::Conversion | TypeUsage::Buffer,
295 );
296 device_props.register_type_usage(
297 ElemType::Float(FloatKind::E3M2),
298 TypeUsage::Conversion | TypeUsage::Buffer,
299 );
300 device_props.register_type_usage(
301 ElemType::Float(FloatKind::UE8M0),
302 TypeUsage::Conversion | TypeUsage::Buffer,
303 );
304
305 if CUDA_VERSION >= 12080 {
306 device_props.features.tma.insert(Tma::SwizzleAtomicity);
307 }
308 }
309
310 device_props.features.memory_reinterpret = true;
311 device_props.features.alignment = true;
312 device_props.features.plane.insert(Plane::Ops);
313 device_props
314 .features
315 .plane
316 .insert(Plane::NonUniformControlFlow);
317
318 register_wmma_features(supported_wmma_combinations, &mut device_props);
319 register_mma_features(supported_mma_combinations, &mut device_props);
320 register_scaled_mma_features(supported_scaled_mma_combinations, &mut device_props);
321
322 let cuda_ctx = CudaContext::new(comp_opts, device_props.clone(), ctx, arch);
323 let logger = Arc::new(ServerLogger::default());
324 let policy = PitchedMemoryLayoutPolicy::new(device_props.memory.alignment as usize);
325 let utilities = ServerUtilities::new(device_props, logger, (), policy);
326
327 CudaServer::new(
328 cuda_ctx,
329 mem_properties,
330 options.memory_config,
331 mem_alignment,
332 device_id,
333 utilities,
334 )
335 }
336
337 fn utilities(&self) -> ServerUtilitiesHandle {
338 self.utilities() as ServerUtilitiesHandle
339 }
340}
341
342pub type CudaCompiler = CppCompiler<CudaDialect<WmmaCompiler>>;
343pub type CudaComputeKernel = ComputeKernel<CudaDialect<WmmaCompiler>>;
344
345fn tensor_cores_per_sm(version: u32) -> Option<u32> {
346 match version {
347 70 | 75 => Some(8), 80 | 86 | 89 | 90 | 91 | 92 | 100 => Some(4), _ => None, }
351}
352
353impl Runtime for CudaRuntime {
354 type Compiler = CudaCompiler;
355 type Server = CudaServer;
356 type Device = CudaDevice;
357
358 fn client(device: &Self::Device) -> ComputeClient<Self> {
359 ComputeClient::load(device)
360 }
361
362 fn name(_client: &ComputeClient<Self>) -> &'static str {
363 "cuda"
364 }
365
366 fn require_array_lengths() -> bool {
367 true
368 }
369
370 fn max_cube_count() -> (u32, u32, u32) {
371 (i32::MAX as u32, u16::MAX as u32, u16::MAX as u32)
372 }
373
374 fn can_read_tensor(shape: &Shape, strides: &Strides) -> bool {
375 has_pitched_row_major_strides(shape, strides)
376 }
377
378 fn target_properties() -> TargetProperties {
379 TargetProperties {
380 mma: MmaProperties {
381 register_size_bits: 32,
382 const_plane_size: 32,
383 register_layout_a: MatrixLayout::RowMajor,
384 register_layout_b: MatrixLayout::ColMajor,
385 register_layout_acc: MatrixLayout::RowMajor,
386 register_duplication_a: 1,
387 register_duplication_b: 1,
388 register_duplication_acc: 1,
389 contiguous_elements: ContiguousElements::new(contiguous_elements_cuda),
390 },
391 }
392 }
393
394 fn enumerate_devices(
395 _: u16,
396 _: &<Self::Server as cubecl_core::server::ComputeServer>::Info,
397 ) -> Vec<cubecl_core::device::DeviceId> {
398 let count = cudarc::driver::CudaContext::device_count().unwrap_or(0) as usize;
399 (0..count)
400 .map(|i| DeviceId {
401 type_id: 0,
402 index_id: i as u16,
403 })
404 .collect()
405 }
406}