1use std::ffi::c_void;
2use std::num::NonZeroUsize;
3use std::sync::{Arc, Mutex};
4
5use cubecl::prelude::{CubeElement, CubePrimitive};
6use cubecl::stream_id::StreamId;
7use cubecl_cuda::CudaRuntime as CubeclCudaRuntime;
8use num_complex::{Complex32, Complex64};
9use num_traits::{One, Zero};
10
11use super::dispatch::{
12 alloc_output, cube_count_for_len, cube_dim_1d, cubecl_buffer, cubecl_view_buffer,
13 cubecl_view_mut_buffer, dtype_mismatch, ensure_resident_on_runtime,
14 ensure_view_mut_resident_on_runtime, ensure_view_resident_on_runtime, launch_nullary_into,
15 prepared_tensor_access, prepared_view_access, prepared_view_mut_access, CubeclPreparedAccess,
16};
17use super::error::{unsupported_dtype, unsupported_operation, workspace_size_overflow};
18use super::ffi::cutensor::{
19 CudaDataType, CutensorComputeDescriptor, CutensorCudaStream, CutensorHandle, CutensorOperator,
20 CutensorWorksizePreference, OperationDescriptor, Plan, PlanPreference, TensorDescriptor,
21};
22use super::interop::cuda_device_ptr_from_addr;
23use super::plan_cache::LruPlanCache;
24use super::{CudaBackend, CudaRuntime};
25use crate::config::DotGeneralConfig;
26use crate::kernels::structural;
27use crate::{col_major_strides, CubeclBuffer, Error, Tensor, TypedTensor};
28use tenferro_tensor::{
29 CacheStats, ContractionScalar, DType, DotGeneralAccumulation, TensorRead, TensorScalar,
30 TensorView, TensorViewMut, TensorWrite, TypedTensorView, TypedTensorViewMut,
31};
32
33const OP: &str = "dot_general";
34const CUDA_ALLOCATION_ALIGNMENT: u32 = 256;
35const DEFAULT_CUTENSOR_PLAN_CACHE_MAX_ENTRIES: usize = 64;
36type CutensorContractionPlanCache = LruPlanCache<CutensorContractionKey, CachedCutensorContraction>;
37type CutensorPlanCacheState = Arc<Mutex<CutensorContractionCacheState>>;
38
39struct CutensorContractionCacheState {
47 workspaces: Box<[Option<Workspace>]>,
51 plans: CutensorContractionPlanCache,
52}
53
54impl CutensorContractionCacheState {
55 fn new(max_entries: NonZeroUsize, stream_slots: usize) -> Self {
56 Self {
57 plans: CutensorContractionPlanCache::new(max_entries),
58 workspaces: (0..stream_slots).map(|_| None).collect(),
59 }
60 }
61
62 fn workspace_bytes(&self) -> u64 {
64 retained_workspace_bytes(&self.workspaces)
65 }
66
67 fn workspace_stats(&self) -> CutensorWorkspaceStats {
74 CutensorWorkspaceStats {
75 retained_entries: self.workspaces.iter().flatten().count(),
76 retained_bytes: self.workspace_bytes(),
77 }
78 }
79
80 fn release_workspaces(&mut self) {
83 for workspace in self.workspaces.iter_mut() {
84 *workspace = None;
85 }
86 }
87}
88
89fn retained_workspace_bytes(workspaces: &[Option<Workspace>]) -> u64 {
90 workspaces.iter().flatten().fold(0_u64, |total, workspace| {
94 total.saturating_add(workspace.size)
95 })
96}
97
98#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
116pub struct CutensorWorkspaceStats {
117 pub retained_entries: usize,
119 pub retained_bytes: u64,
121}
122
123#[derive(Clone, Copy, Debug, PartialEq, Eq)]
125pub(super) enum WorkspacePlan {
126 Reuse,
128 Retain(u64),
130 Temporary(u64),
133}
134
135pub(super) fn plan_workspace(
146 requested: u64,
147 current_capacity: u64,
148 retained_total: u64,
149 limit: u64,
150) -> WorkspacePlan {
151 if current_capacity >= requested {
152 return WorkspacePlan::Reuse;
153 }
154 let headroom = limit.saturating_sub(retained_total.saturating_sub(current_capacity));
158 match shared_workspace_capacity(requested) {
159 Some(rounded) if rounded <= headroom => WorkspacePlan::Retain(rounded),
160 _ if requested <= headroom => WorkspacePlan::Retain(requested),
163 _ => WorkspacePlan::Temporary(requested),
164 }
165}
166
167trait CutensorScalar: CubeElement + TensorScalar + CubePrimitive + Clone + One + Zero {
168 const DATA_TYPE: CudaDataType;
169 const DTYPE: DType;
170 const IS_COMPLEX: bool;
171
172 fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor;
173
174 fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>>;
176
177 fn unwrap_view<'a, 'b>(view: &'a TensorView<'b>) -> Option<&'a TypedTensorView<'b, Self>>;
179 fn unwrap_view_mut<'a, 'b>(
180 view: &'a mut TensorViewMut<'b>,
181 ) -> Option<&'a mut TypedTensorViewMut<'b, Self>>;
182 fn unwrap_tensor_mut(tensor: &mut Tensor) -> Option<&mut TypedTensor<Self>>;
183
184 fn launch_scale_in_place(
186 client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
187 count: cubecl::prelude::CubeCount,
188 dim: cubecl::prelude::CubeDim,
189 out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
190 factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
191 );
192}
193
194macro_rules! cutensor_variant_accessors {
196 ($variant:ident) => {
197 fn unwrap_view<'a, 'b>(view: &'a TensorView<'b>) -> Option<&'a TypedTensorView<'b, Self>> {
198 match view {
199 TensorView::$variant(view) => Some(view),
200 _ => None,
201 }
202 }
203
204 fn unwrap_view_mut<'a, 'b>(
205 view: &'a mut TensorViewMut<'b>,
206 ) -> Option<&'a mut TypedTensorViewMut<'b, Self>> {
207 match view {
208 TensorViewMut::$variant(view) => Some(view),
209 _ => None,
210 }
211 }
212
213 fn unwrap_tensor_mut(tensor: &mut Tensor) -> Option<&mut TypedTensor<Self>> {
214 tensor.as_typed_mut::<Self>()
215 }
216 };
217}
218
219impl CutensorScalar for f32 {
220 cutensor_variant_accessors!(F32);
221
222 const DATA_TYPE: CudaDataType = CudaDataType::R32F;
223 const DTYPE: DType = DType::F32;
224 const IS_COMPLEX: bool = false;
225
226 fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
227 handle.compute_desc_32f()
228 }
229 fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
230 tensor.as_typed::<Self>()
231 }
232
233 fn launch_scale_in_place(
234 client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
235 count: cubecl::prelude::CubeCount,
236 dim: cubecl::prelude::CubeDim,
237 out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
238 factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
239 ) {
240 unsafe {
242 structural::scale_in_place_float_kernel::launch_unchecked::<f32, CubeclCudaRuntime>(
243 client, count, dim, out, factor,
244 );
245 }
246 }
247}
248
249impl CutensorScalar for f64 {
250 cutensor_variant_accessors!(F64);
251
252 const DATA_TYPE: CudaDataType = CudaDataType::R64F;
253 const DTYPE: DType = DType::F64;
254 const IS_COMPLEX: bool = false;
255
256 fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
257 handle.compute_desc_64f()
258 }
259 fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
260 tensor.as_typed::<Self>()
261 }
262
263 fn launch_scale_in_place(
264 client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
265 count: cubecl::prelude::CubeCount,
266 dim: cubecl::prelude::CubeDim,
267 out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
268 factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
269 ) {
270 unsafe {
272 structural::scale_in_place_float_kernel::launch_unchecked::<f64, CubeclCudaRuntime>(
273 client, count, dim, out, factor,
274 );
275 }
276 }
277}
278
279impl CutensorScalar for Complex32 {
280 cutensor_variant_accessors!(C32);
281
282 const DATA_TYPE: CudaDataType = CudaDataType::C32F;
283 const DTYPE: DType = DType::C32;
284 const IS_COMPLEX: bool = true;
285
286 fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
287 handle.compute_desc_32f()
288 }
289 fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
290 tensor.as_typed::<Self>()
291 }
292
293 fn launch_scale_in_place(
294 client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
295 count: cubecl::prelude::CubeCount,
296 dim: cubecl::prelude::CubeDim,
297 out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
298 factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
299 ) {
300 unsafe {
302 structural::scale_in_place_complex_kernel::launch_unchecked::<
303 Complex32,
304 CubeclCudaRuntime,
305 >(client, count, dim, out, factor);
306 }
307 }
308}
309
310impl CutensorScalar for Complex64 {
311 cutensor_variant_accessors!(C64);
312
313 const DATA_TYPE: CudaDataType = CudaDataType::C64F;
314 const DTYPE: DType = DType::C64;
315 const IS_COMPLEX: bool = true;
316
317 fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
318 handle.compute_desc_64f()
319 }
320 fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
321 tensor.as_typed::<Self>()
322 }
323
324 fn launch_scale_in_place(
325 client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
326 count: cubecl::prelude::CubeCount,
327 dim: cubecl::prelude::CubeDim,
328 out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
329 factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
330 ) {
331 unsafe {
333 structural::scale_in_place_complex_kernel::launch_unchecked::<
334 Complex64,
335 CubeclCudaRuntime,
336 >(client, count, dim, out, factor);
337 }
338 }
339}
340
341struct DotGeneralLayout {
342 lhs_modes: Vec<i32>,
343 rhs_modes: Vec<i32>,
344 output_modes: Vec<i32>,
345 output_shape: Vec<usize>,
346 lhs_extents: Vec<i64>,
347 rhs_extents: Vec<i64>,
348 output_extents: Vec<i64>,
349 lhs_strides: Vec<i64>,
350 rhs_strides: Vec<i64>,
351 output_strides: Vec<i64>,
352 contracting_elements: usize,
353}
354
355struct Workspace {
356 _handle: Option<cubecl_runtime::server::Handle>,
359 ptr: *mut c_void,
360 size: u64,
361 runtime: Option<CudaRuntime>,
362 stream: u64,
363}
364
365impl Workspace {
366 fn none() -> Self {
367 Self {
368 _handle: None,
369 ptr: std::ptr::null_mut(),
370 size: 0,
371 runtime: None,
372 stream: 0,
373 }
374 }
375}
376
377impl Drop for Workspace {
378 fn drop(&mut self) {
379 let (Some(runtime), Some(handle)) = (self.runtime.as_ref(), self._handle.take()) else {
380 return;
381 };
382 let mut retirements = runtime
386 .workspace_retirements()
387 .lock()
388 .unwrap_or_else(|error| error.into_inner());
389 retirements.retire(runtime.state(), self.stream, handle);
390 }
391}
392
393unsafe impl Send for Workspace {}
397
398#[derive(Clone, Debug, PartialEq, Eq, Hash)]
399struct CutensorOperandLayoutKey {
400 extents: Vec<i64>,
401 strides: Vec<i64>,
402 modes: Vec<i32>,
403}
404
405impl CutensorOperandLayoutKey {
406 fn new(extents: &[i64], strides: &[i64], modes: &[i32]) -> Self {
407 Self {
408 extents: extents.to_vec(),
409 strides: strides.to_vec(),
410 modes: modes.to_vec(),
411 }
412 }
413
414 fn retained_bytes(&self) -> usize {
415 std::mem::size_of::<Self>()
416 .saturating_add(self.extents.capacity() * std::mem::size_of::<i64>())
417 .saturating_add(self.strides.capacity() * std::mem::size_of::<i64>())
418 .saturating_add(self.modes.capacity() * std::mem::size_of::<i32>())
419 }
420}
421
422#[derive(Clone, Debug, PartialEq, Eq, Hash)]
423struct CutensorContractionKey {
424 dtype: DType,
425 lhs: CutensorOperandLayoutKey,
426 rhs: CutensorOperandLayoutKey,
427 output: CutensorOperandLayoutKey,
428 lhs_alignment_requirement: u32,
429 rhs_alignment_requirement: u32,
430 output_alignment_requirement: u32,
431 lhs_op: CutensorOperator,
432 rhs_op: CutensorOperator,
433 workspace_preference: CutensorWorksizePreference,
434}
435
436impl CutensorContractionKey {
437 fn from_spec<T: CutensorScalar>(spec: &CutensorContractionSpec<'_>) -> Self {
438 Self {
439 dtype: T::DTYPE,
440 lhs: CutensorOperandLayoutKey::new(
441 &spec.layout.lhs_extents,
442 spec.lhs_strides,
443 &spec.layout.lhs_modes,
444 ),
445 rhs: CutensorOperandLayoutKey::new(
446 &spec.layout.rhs_extents,
447 spec.rhs_strides,
448 &spec.layout.rhs_modes,
449 ),
450 output: CutensorOperandLayoutKey::new(
451 &spec.layout.output_extents,
452 spec.output_strides,
453 &spec.layout.output_modes,
454 ),
455 lhs_alignment_requirement: spec.lhs_alignment_requirement,
456 rhs_alignment_requirement: spec.rhs_alignment_requirement,
457 output_alignment_requirement: spec.output_alignment_requirement,
458 lhs_op: cutensor_conj_op::<T>(spec.lhs_conj),
459 rhs_op: cutensor_conj_op::<T>(spec.rhs_conj),
460 workspace_preference: spec.workspace_preference,
461 }
462 }
463
464 fn retained_bytes(&self) -> usize {
465 std::mem::size_of::<Self>()
466 .saturating_add(self.lhs.retained_bytes())
467 .saturating_add(self.rhs.retained_bytes())
468 .saturating_add(self.output.retained_bytes())
469 }
470}
471
472struct CutensorContractionSpec<'a> {
473 layout: &'a DotGeneralLayout,
474 lhs_strides: &'a [i64],
475 rhs_strides: &'a [i64],
476 output_strides: &'a [i64],
477 lhs_alignment_requirement: u32,
478 rhs_alignment_requirement: u32,
479 output_alignment_requirement: u32,
480 lhs_conj: bool,
481 rhs_conj: bool,
482 workspace_preference: CutensorWorksizePreference,
483}
484
485struct CachedCutensorContraction {
486 workspace_size: u64,
490 plan: Plan,
492 _plan_preference: PlanPreference,
493 _operation_descriptor: OperationDescriptor,
494 _output_descriptor: TensorDescriptor,
495 _rhs_descriptor: TensorDescriptor,
496 _lhs_descriptor: TensorDescriptor,
497}
498
499unsafe impl Send for CachedCutensorContraction {}
503
504impl CachedCutensorContraction {
505 fn new<T>(cutensor: &CutensorHandle, spec: &CutensorContractionSpec<'_>) -> crate::Result<Self>
506 where
507 T: CutensorScalar,
508 {
509 let desc_a = TensorDescriptor::new(
510 cutensor,
511 &spec.layout.lhs_extents,
512 spec.lhs_strides,
513 T::DATA_TYPE,
514 spec.lhs_alignment_requirement,
515 OP,
516 )?;
517 let desc_b = TensorDescriptor::new(
518 cutensor,
519 &spec.layout.rhs_extents,
520 spec.rhs_strides,
521 T::DATA_TYPE,
522 spec.rhs_alignment_requirement,
523 OP,
524 )?;
525 let desc_out = TensorDescriptor::new(
526 cutensor,
527 &spec.layout.output_extents,
528 spec.output_strides,
529 T::DATA_TYPE,
530 spec.output_alignment_requirement,
531 OP,
532 )?;
533 let op_desc = OperationDescriptor::new_contraction_with_ops(
534 cutensor,
535 &desc_a,
536 &spec.layout.lhs_modes,
537 cutensor_conj_op::<T>(spec.lhs_conj),
538 &desc_b,
539 &spec.layout.rhs_modes,
540 cutensor_conj_op::<T>(spec.rhs_conj),
541 &desc_out,
542 &spec.layout.output_modes,
543 &desc_out,
544 &spec.layout.output_modes,
545 T::compute_descriptor(cutensor),
546 OP,
547 )?;
548 let pref = PlanPreference::new_default(cutensor, OP)?;
549 let workspace_size =
550 cutensor.estimate_workspace_size(&op_desc, &pref, spec.workspace_preference, OP)?;
551 let plan = Plan::new(cutensor, &op_desc, &pref, workspace_size, OP)?;
552 Ok(Self {
553 workspace_size,
554 plan,
555 _plan_preference: pref,
556 _operation_descriptor: op_desc,
557 _output_descriptor: desc_out,
558 _rhs_descriptor: desc_b,
559 _lhs_descriptor: desc_a,
560 })
561 }
562
563 fn retained_bytes(&self) -> usize {
564 std::mem::size_of::<Self>()
565 }
566}
567
568fn spec_hash<T: CutensorScalar>(spec: &CutensorContractionSpec<'_>) -> u64 {
573 use std::hash::{Hash, Hasher};
574 let mut hasher = std::collections::hash_map::DefaultHasher::new();
575 T::DTYPE.hash(&mut hasher);
576 spec.layout.lhs_extents.hash(&mut hasher);
577 spec.lhs_strides.hash(&mut hasher);
578 spec.layout.lhs_modes.hash(&mut hasher);
579 spec.layout.rhs_extents.hash(&mut hasher);
580 spec.rhs_strides.hash(&mut hasher);
581 spec.layout.rhs_modes.hash(&mut hasher);
582 spec.layout.output_extents.hash(&mut hasher);
583 spec.output_strides.hash(&mut hasher);
584 spec.layout.output_modes.hash(&mut hasher);
585 spec.lhs_alignment_requirement.hash(&mut hasher);
586 spec.rhs_alignment_requirement.hash(&mut hasher);
587 spec.output_alignment_requirement.hash(&mut hasher);
588 cutensor_conj_op::<T>(spec.lhs_conj).hash(&mut hasher);
589 cutensor_conj_op::<T>(spec.rhs_conj).hash(&mut hasher);
590 spec.workspace_preference.hash(&mut hasher);
591 hasher.finish()
592}
593
594fn key_matches_spec<T: CutensorScalar>(
596 key: &CutensorContractionKey,
597 spec: &CutensorContractionSpec<'_>,
598) -> bool {
599 key.dtype == T::DTYPE
600 && key.lhs.extents == spec.layout.lhs_extents
601 && key.lhs.strides == spec.lhs_strides
602 && key.lhs.modes == spec.layout.lhs_modes
603 && key.rhs.extents == spec.layout.rhs_extents
604 && key.rhs.strides == spec.rhs_strides
605 && key.rhs.modes == spec.layout.rhs_modes
606 && key.output.extents == spec.layout.output_extents
607 && key.output.strides == spec.output_strides
608 && key.output.modes == spec.layout.output_modes
609 && key.lhs_alignment_requirement == spec.lhs_alignment_requirement
610 && key.rhs_alignment_requirement == spec.rhs_alignment_requirement
611 && key.output_alignment_requirement == spec.output_alignment_requirement
612 && key.lhs_op == cutensor_conj_op::<T>(spec.lhs_conj)
613 && key.rhs_op == cutensor_conj_op::<T>(spec.rhs_conj)
614 && key.workspace_preference == spec.workspace_preference
615}
616
617fn gemm_pair_operands<'a, T: TensorScalar>(
619 op: &'static str,
620 lhs: &'a Tensor,
621 rhs: &'a Tensor,
622) -> crate::Result<(&'a TypedTensor<T>, &'a TypedTensor<T>)> {
623 let lhs_t = lhs
624 .as_typed::<T>()
625 .ok_or_else(|| dtype_mismatch(op, lhs, rhs))?;
626 let rhs_t = rhs
627 .as_typed::<T>()
628 .ok_or_else(|| dtype_mismatch(op, lhs, rhs))?;
629 Ok((lhs_t, rhs_t))
630}
631
632pub(super) fn dot_general_with_conj(
633 backend: &CudaBackend,
634 lhs: &Tensor,
635 rhs: &Tensor,
636 config: &DotGeneralConfig,
637 lhs_conj: bool,
638 rhs_conj: bool,
639) -> crate::Result<Tensor> {
640 match (lhs.dtype(), rhs.dtype()) {
641 (DType::F32, DType::F32) => {
642 let (lhs, rhs) = gemm_pair_operands::<f32>(OP, lhs, rhs)?;
643 dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
644 .map(Tensor::from_typed::<f32>)
645 }
646 (DType::F64, DType::F64) => {
647 let (lhs, rhs) = gemm_pair_operands::<f64>(OP, lhs, rhs)?;
648 dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
649 .map(Tensor::from_typed::<f64>)
650 }
651 (DType::C32, DType::C32) => {
652 let (lhs, rhs) = gemm_pair_operands::<Complex32>(OP, lhs, rhs)?;
653 dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
654 .map(Tensor::from_typed::<Complex32>)
655 }
656 (DType::C64, DType::C64) => {
657 let (lhs, rhs) = gemm_pair_operands::<Complex64>(OP, lhs, rhs)?;
658 dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
659 .map(Tensor::from_typed::<Complex64>)
660 }
661 _ => Err(dtype_mismatch(OP, lhs, rhs)),
662 }
663}
664
665trait FromContractionScalar: Sized {
668 fn from_contraction_scalar(value: ContractionScalar) -> crate::Result<Self>;
669}
670
671macro_rules! impl_from_contraction_scalar {
672 ($ty:ty, $variant:ident) => {
673 impl FromContractionScalar for $ty {
674 fn from_contraction_scalar(value: ContractionScalar) -> crate::Result<Self> {
675 match value {
676 ContractionScalar::$variant(value) => Ok(value),
677 other => Err(Error::dtype_mismatch(
678 OP,
679 <$ty as tenferro_tensor::TensorScalar>::dtype(),
680 other.dtype(),
681 )),
682 }
683 }
684 }
685 };
686}
687
688impl_from_contraction_scalar!(f32, F32);
689impl_from_contraction_scalar!(f64, F64);
690impl_from_contraction_scalar!(Complex32, C32);
691impl_from_contraction_scalar!(Complex64, C64);
692
693enum ReadOperand<'a, 'b, T> {
696 Owned(&'a TypedTensor<T>),
697 View(&'a TypedTensorView<'b, T>),
698}
699
700impl<T: 'static> ReadOperand<'_, '_, T> {
701 fn shape(&self) -> &[usize] {
702 match self {
703 Self::Owned(tensor) => tensor.shape(),
704 Self::View(view) => view.shape(),
705 }
706 }
707
708 fn handle(&self) -> crate::Result<&cubecl_runtime::server::Handle> {
709 match self {
710 Self::Owned(tensor) => Ok(cubecl_buffer(tensor, OP)?.handle()),
711 Self::View(view) => Ok(cubecl_view_buffer(view, OP)?.handle()),
712 }
713 }
714}
715
716fn read_operand_alignment_requirement<T: CutensorScalar>(operand: &ReadOperand<'_, '_, T>) -> u32 {
717 match operand {
718 ReadOperand::Owned(_) => CUDA_ALLOCATION_ALIGNMENT,
719 ReadOperand::View(_) => view_descriptor_alignment_requirement::<T>(),
720 }
721}
722
723enum WriteOperand<'a, 'b, T> {
725 Owned(&'a mut TypedTensor<T>),
726 View(&'a mut TypedTensorViewMut<'b, T>),
727}
728
729impl<T: 'static> WriteOperand<'_, '_, T> {
730 fn shape(&self) -> &[usize] {
731 match self {
732 Self::Owned(tensor) => tensor.shape(),
733 Self::View(view) => view.shape(),
734 }
735 }
736
737 fn n_elements(&self) -> usize {
738 match self {
739 Self::Owned(tensor) => tensor.n_elements(),
740 Self::View(view) => view.n_elements(),
741 }
742 }
743
744 fn handle(&self) -> crate::Result<&cubecl_runtime::server::Handle> {
745 match self {
746 Self::Owned(tensor) => Ok(cubecl_buffer(tensor, OP)?.handle()),
747 Self::View(view) => Ok(cubecl_view_mut_buffer(view, OP)?.handle()),
748 }
749 }
750}
751
752fn cross_stream_handles<'a>(
753 rt: &CudaRuntime,
754 handles: impl IntoIterator<Item = &'a cubecl_runtime::server::Handle>,
755) -> Vec<cubecl_runtime::server::Handle> {
756 handles
757 .into_iter()
758 .filter(|handle| !rt.is_current_stream_slot(handle))
759 .cloned()
760 .collect()
761}
762
763fn write_operand_alignment_requirement<T: CutensorScalar>(
764 operand: &WriteOperand<'_, '_, T>,
765) -> u32 {
766 match operand {
767 WriteOperand::Owned(_) => CUDA_ALLOCATION_ALIGNMENT,
768 WriteOperand::View(_) => view_descriptor_alignment_requirement::<T>(),
769 }
770}
771
772fn view_descriptor_alignment_requirement<T: CutensorScalar>() -> u32 {
773 u32::try_from(std::mem::size_of::<T>()).unwrap_or(CUDA_ALLOCATION_ALIGNMENT)
774}
775
776struct ResolvedOperand<'a> {
780 ptr: *mut c_void,
781 strides: std::borrow::Cow<'a, [i64]>,
782 alignment: u32,
783}
784
785fn read_operand<'a, 'b, T: CutensorScalar>(
786 read: &'a TensorRead<'b>,
787) -> Option<ReadOperand<'a, 'b, T>> {
788 match read {
789 TensorRead::Tensor(tensor) => T::unwrap_tensor(tensor).map(ReadOperand::Owned),
790 TensorRead::View(view) => T::unwrap_view(view).map(ReadOperand::View),
791 }
792}
793
794fn write_operand<'a, 'b, T: CutensorScalar>(
795 write: &'a mut TensorWrite<'b>,
796) -> Option<WriteOperand<'a, 'b, T>> {
797 match write {
798 TensorWrite::Tensor(tensor) => T::unwrap_tensor_mut(tensor).map(WriteOperand::Owned),
799 TensorWrite::View(view) => T::unwrap_view_mut(view).map(WriteOperand::View),
800 }
801}
802
803pub(super) fn dot_general_read_allocating(
820 backend: &mut CudaBackend,
821 lhs: TensorRead<'_>,
822 rhs: TensorRead<'_>,
823 config: &DotGeneralConfig,
824 lhs_conj: bool,
825 rhs_conj: bool,
826) -> crate::Result<Tensor> {
827 if let (Some(lhs_owned), Some(rhs_owned)) = (lhs.as_tensor(), rhs.as_tensor()) {
828 return dot_general_with_conj(backend, lhs_owned, rhs_owned, config, lhs_conj, rhs_conj);
829 }
830 let dtype = lhs.dtype();
831 let shape =
832 tenferro_tensor::backend::dot_general_output_shape(lhs.shape(), rhs.shape(), config, OP)?;
833 let mut out = match dtype {
834 DType::F32 => Tensor::from_typed::<f32>(alloc_output::<f32>(backend.runtime(), &shape)?),
835 DType::F64 => Tensor::from_typed::<f64>(alloc_output::<f64>(backend.runtime(), &shape)?),
836 DType::C32 => {
837 Tensor::from_typed::<Complex32>(alloc_output::<Complex32>(backend.runtime(), &shape)?)
838 }
839 DType::C64 => {
840 Tensor::from_typed::<Complex64>(alloc_output::<Complex64>(backend.runtime(), &shape)?)
841 }
842 dtype => return Err(unsupported_dtype(OP, dtype)),
843 };
844 let accumulation = DotGeneralAccumulation {
845 lhs_conj,
846 rhs_conj,
847 ..DotGeneralAccumulation::overwrite(dtype)?
848 };
849 {
850 let mut out_write = TensorWrite::from_tensor(&mut out);
851 dot_general_read_into_accum(backend, &lhs, &rhs, config, accumulation, &mut out_write)?;
852 }
853 Ok(out)
854}
855
856pub(super) fn dot_general_read_into_accum(
857 backend: &CudaBackend,
858 lhs: &TensorRead<'_>,
859 rhs: &TensorRead<'_>,
860 config: &DotGeneralConfig,
861 accumulation: DotGeneralAccumulation,
862 out: &mut TensorWrite<'_>,
863) -> crate::Result<()> {
864 match lhs.dtype() {
865 DType::F32 => accum_erased::<f32>(backend, lhs, rhs, config, accumulation, out),
866 DType::F64 => accum_erased::<f64>(backend, lhs, rhs, config, accumulation, out),
867 DType::C32 => accum_erased::<Complex32>(backend, lhs, rhs, config, accumulation, out),
868 DType::C64 => accum_erased::<Complex64>(backend, lhs, rhs, config, accumulation, out),
869 dtype => Err(unsupported_dtype(OP, dtype)),
870 }
871}
872
873fn accum_erased<T>(
874 backend: &CudaBackend,
875 lhs: &TensorRead<'_>,
876 rhs: &TensorRead<'_>,
877 config: &DotGeneralConfig,
878 accumulation: DotGeneralAccumulation,
879 out: &mut TensorWrite<'_>,
880) -> crate::Result<()>
881where
882 T: CutensorScalar + FromContractionScalar + PartialEq + tenferro_tensor::TensorScalar,
883{
884 let (lhs_dtype, rhs_dtype, out_dtype) = (lhs.dtype(), rhs.dtype(), out.dtype());
885 let (Some(lhs), Some(rhs), Some(out)) = (
886 read_operand::<T>(lhs),
887 read_operand::<T>(rhs),
888 write_operand::<T>(out),
889 ) else {
890 let (expected, actual) = if lhs_dtype != rhs_dtype {
891 (lhs_dtype, rhs_dtype)
892 } else {
893 (lhs_dtype, out_dtype)
894 };
895 return Err(Error::dtype_mismatch(OP, expected, actual));
896 };
897 dot_general_typed_into_accum(
898 backend,
899 lhs,
900 rhs,
901 config,
902 accumulation.lhs_conj,
903 accumulation.rhs_conj,
904 T::from_contraction_scalar(accumulation.alpha)?,
905 T::from_contraction_scalar(accumulation.beta)?,
906 out,
907 )
908}
909
910#[allow(clippy::too_many_arguments)]
911fn dot_general_typed_into_accum<T>(
912 backend: &CudaBackend,
913 lhs: ReadOperand<'_, '_, T>,
914 rhs: ReadOperand<'_, '_, T>,
915 config: &DotGeneralConfig,
916 lhs_conj: bool,
917 rhs_conj: bool,
918 alpha: T,
919 beta: T,
920 mut out: WriteOperand<'_, '_, T>,
921) -> crate::Result<()>
922where
923 T: CutensorScalar + PartialEq + tenferro_tensor::TensorScalar,
924{
925 backend.runtime().set_current_cuda_context(OP)?;
926 validate_dot_general(lhs.shape(), rhs.shape(), config)?;
927 let layout = build_layout(lhs.shape(), rhs.shape(), config)?;
928 if out.shape() != layout.output_shape.as_slice() {
929 return Err(Error::shape_mismatch(
930 OP,
931 out.shape().to_vec(),
932 layout.output_shape.clone(),
933 ));
934 }
935 let cross_stream_handles = cross_stream_handles(
936 backend.runtime(),
937 [lhs.handle()?, rhs.handle()?, out.handle()?],
938 );
939 let lhs_res = resolve_read_operand(backend.runtime(), &lhs, &layout.lhs_strides)?;
942 let rhs_res = resolve_read_operand(backend.runtime(), &rhs, &layout.rhs_strides)?;
943 let out_res = resolve_write_operand(backend.runtime(), &mut out, &layout.output_strides)?;
944 if out.n_elements() == 0 {
945 return Ok(());
946 }
947 if layout.contracting_elements == 0 {
948 return match out {
950 WriteOperand::Owned(tensor) => scale_in_place(backend.runtime(), tensor, beta),
951 WriteOperand::View(_) => {
952 if beta == T::one() {
953 Ok(())
954 } else {
955 Err(unsupported_operation(
958 OP,
959 "zero-sized contraction with beta != 1 is not supported for borrowed view outputs",
960 ))
961 }
962 }
963 };
964 }
965
966 let stream = raw_stream(backend.runtime())?;
967 let spec = CutensorContractionSpec {
968 layout: &layout,
969 lhs_strides: &lhs_res.strides,
970 rhs_strides: &rhs_res.strides,
971 output_strides: &out_res.strides,
972 lhs_alignment_requirement: read_operand_alignment_requirement(&lhs),
973 rhs_alignment_requirement: read_operand_alignment_requirement(&rhs),
974 output_alignment_requirement: write_operand_alignment_requirement(&out),
975 lhs_conj,
976 rhs_conj,
977 workspace_preference: CutensorWorksizePreference::Default,
978 };
979 validate_descriptor_alignment(lhs_res.alignment, spec.lhs_alignment_requirement, "lhs")?;
980 validate_descriptor_alignment(rhs_res.alignment, spec.rhs_alignment_requirement, "rhs")?;
981 validate_descriptor_alignment(out_res.alignment, spec.output_alignment_requirement, "out")?;
982 cached_cutensor_contraction::<T, _>(
987 backend,
988 &spec,
989 cross_stream_handles,
990 |cutensor, plan, workspace| unsafe {
991 cutensor.contract(
992 plan,
993 &alpha as *const T as *const c_void,
994 lhs_res.ptr as *const c_void,
995 rhs_res.ptr as *const c_void,
996 &beta as *const T as *const c_void,
997 out_res.ptr as *const c_void,
998 out_res.ptr,
999 workspace.ptr,
1000 workspace.size,
1001 stream,
1002 OP,
1003 )
1004 },
1005 )
1006}
1007
1008fn resolve_read_operand<'a, T>(
1009 rt: &CudaRuntime,
1010 operand: &ReadOperand<'_, '_, T>,
1011 compact_strides: &'a [i64],
1012) -> crate::Result<ResolvedOperand<'a>>
1013where
1014 T: CutensorScalar + 'static,
1015{
1016 match operand {
1017 ReadOperand::Owned(tensor) => Ok(ResolvedOperand {
1018 ptr: typed_device_ptr(rt, tensor, OP)?,
1019 strides: std::borrow::Cow::Borrowed(compact_strides),
1020 alignment: CUDA_ALLOCATION_ALIGNMENT,
1021 }),
1022 ReadOperand::View(view) => {
1023 ensure_view_resident_on_runtime(rt, view, OP)?;
1024 let prepared = prepared_view_access(view, OP)?;
1025 let base = memoized_device_addr(rt, cubecl_view_buffer(view, OP)?, prepared, OP)?;
1028 resolve_prepared_device_region::<T>(base, view.strides(), view.offset())
1029 }
1030 }
1031}
1032
1033fn resolve_write_operand<'a, T>(
1034 rt: &CudaRuntime,
1035 operand: &mut WriteOperand<'_, '_, T>,
1036 compact_strides: &'a [i64],
1037) -> crate::Result<ResolvedOperand<'a>>
1038where
1039 T: CutensorScalar + 'static,
1040{
1041 match operand {
1042 WriteOperand::Owned(tensor) => Ok(ResolvedOperand {
1043 ptr: write_device_ptr(rt, tensor, OP)?,
1044 strides: std::borrow::Cow::Borrowed(compact_strides),
1045 alignment: CUDA_ALLOCATION_ALIGNMENT,
1046 }),
1047 WriteOperand::View(view) => {
1048 ensure_view_mut_resident_on_runtime(rt, view, OP)?;
1049 let prepared = prepared_view_mut_access(view, OP)?;
1050 let base = rt
1053 .client()
1054 .get_resource(prepared.into_handle())
1055 .map_err(|err| Error::backend_source(OP, err))?
1056 .resource()
1057 .ptr;
1058 resolve_prepared_device_region::<T>(base, view.strides(), view.offset())
1059 }
1060 }
1061}
1062
1063fn resolve_prepared_device_region<T: CutensorScalar + 'static>(
1068 base_addr: u64,
1069 strides: &[isize],
1070 offset: isize,
1071) -> crate::Result<ResolvedOperand<'static>> {
1072 let mut strides_i64 = Vec::with_capacity(strides.len());
1073 for &stride in strides {
1074 if stride < 0 {
1075 return Err(Error::invalid_argument(
1076 OP,
1077 "layout",
1078 format!(
1079 "cuTENSOR dot-general accumulation requires nonnegative view strides, got {strides:?}; canonicalize the view on device first"
1080 ),
1081 ));
1082 }
1083 strides_i64.push(stride as i64);
1084 }
1085 let offset = usize::try_from(offset)
1086 .map_err(|_| Error::invalid_argument(OP, "layout", "view offset must be nonnegative"))?;
1087 let offset_bytes = offset
1088 .checked_mul(std::mem::size_of::<T>())
1089 .ok_or_else(|| Error::invalid_argument(OP, "layout", "view byte offset overflows"))?;
1090 let addr = base_addr
1091 .checked_add(offset_bytes as u64)
1092 .ok_or_else(|| Error::invalid_argument(OP, "layout", "view device address overflows"))?;
1093 Ok(ResolvedOperand {
1098 ptr: cuda_device_ptr_from_addr(addr, OP)?,
1099 strides: std::borrow::Cow::Owned(strides_i64),
1100 alignment: view_descriptor_alignment_requirement::<T>(),
1101 })
1102}
1103
1104fn scale_in_place<T>(rt: &CudaRuntime, out: &mut TypedTensor<T>, beta: T) -> crate::Result<()>
1108where
1109 T: CutensorScalar + PartialEq + tenferro_tensor::TensorScalar,
1110{
1111 if beta == T::one() {
1112 return Ok(());
1113 }
1114 if beta == T::zero() {
1115 return launch_nullary_into(
1116 rt,
1117 out,
1118 OP,
1119 cube_count_for_len(out.n_elements())?,
1120 cube_dim_1d(),
1121 |client, count, dim, out| unsafe {
1122 structural::fill_zero_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1123 client, count, dim, out,
1124 );
1125 },
1126 );
1127 }
1128 super::interop::scale_typed_tensor_for_op(rt, out, beta, OP, T::launch_scale_in_place)
1129}
1130
1131fn dot_general_typed_with_conj<T>(
1132 backend: &CudaBackend,
1133 lhs: &TypedTensor<T>,
1134 rhs: &TypedTensor<T>,
1135 config: &DotGeneralConfig,
1136 lhs_conj: bool,
1137 rhs_conj: bool,
1138) -> crate::Result<TypedTensor<T>>
1139where
1140 T: CutensorScalar,
1141{
1142 backend.runtime().set_current_cuda_context(OP)?;
1143 validate_dot_general(lhs.shape(), rhs.shape(), config)?;
1144 let layout = build_layout(lhs.shape(), rhs.shape(), config)?;
1145 let output = alloc_output::<T>(backend.runtime(), &layout.output_shape)?;
1146 if output.n_elements() == 0 {
1147 return Ok(output);
1148 }
1149 if layout.contracting_elements == 0 {
1150 launch_nullary_into(
1153 backend.runtime(),
1154 &output,
1155 OP,
1156 cube_count_for_len(output.n_elements())?,
1157 cube_dim_1d(),
1158 |client, count, dim, out| unsafe {
1159 structural::fill_zero_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1160 client, count, dim, out,
1161 );
1162 },
1163 )?;
1164 return Ok(output);
1165 }
1166
1167 let lhs_ptr = typed_device_ptr(backend.runtime(), lhs, OP)?;
1168 let rhs_ptr = typed_device_ptr(backend.runtime(), rhs, OP)?;
1169 let output_ptr = typed_device_ptr(backend.runtime(), &output, OP)?;
1170
1171 let alpha = T::one();
1172 let beta = T::zero();
1173 let stream = raw_stream(backend.runtime())?;
1174 let spec = CutensorContractionSpec {
1175 layout: &layout,
1176 lhs_strides: &layout.lhs_strides,
1177 rhs_strides: &layout.rhs_strides,
1178 output_strides: &layout.output_strides,
1179 lhs_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1180 rhs_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1181 output_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1182 lhs_conj,
1183 rhs_conj,
1184 workspace_preference: CutensorWorksizePreference::Default,
1185 };
1186 let cross_stream_handles = cross_stream_handles(
1190 backend.runtime(),
1191 [
1192 cubecl_buffer(lhs, OP)?.handle(),
1193 cubecl_buffer(rhs, OP)?.handle(),
1194 cubecl_buffer(&output, OP)?.handle(),
1195 ],
1196 );
1197 cached_cutensor_contraction::<T, _>(
1198 backend,
1199 &spec,
1200 cross_stream_handles,
1201 |cutensor, plan, workspace| unsafe {
1202 cutensor.contract(
1203 plan,
1204 &alpha as *const T as *const c_void,
1205 lhs_ptr as *const c_void,
1206 rhs_ptr as *const c_void,
1207 &beta as *const T as *const c_void,
1208 output_ptr as *const c_void,
1209 output_ptr,
1210 workspace.ptr,
1211 workspace.size,
1212 stream,
1213 OP,
1214 )
1215 },
1216 )?;
1217
1218 Ok(output)
1219}
1220
1221fn cutensor_conj_op<T: CutensorScalar>(conj: bool) -> CutensorOperator {
1222 if conj && T::IS_COMPLEX {
1223 CutensorOperator::Conj
1224 } else {
1225 CutensorOperator::Identity
1226 }
1227}
1228
1229fn default_cutensor_plan_cache_max_entries() -> NonZeroUsize {
1230 NonZeroUsize::new(DEFAULT_CUTENSOR_PLAN_CACHE_MAX_ENTRIES).unwrap_or(NonZeroUsize::MIN)
1231}
1232
1233fn new_cutensor_plan_cache_state(
1234 max_entries: NonZeroUsize,
1235 stream_slots: usize,
1236) -> CutensorPlanCacheState {
1237 Arc::new(Mutex::new(CutensorContractionCacheState::new(
1238 max_entries,
1239 stream_slots,
1240 )))
1241}
1242
1243fn get_or_init_cutensor_plan_cache(backend: &CudaBackend) -> crate::Result<CutensorPlanCacheState> {
1244 let guard = backend
1245 .cuda_extension_cache()
1246 .get_or_try_init::<CutensorPlanCacheState>(|| {
1247 Ok(new_cutensor_plan_cache_state(
1248 default_cutensor_plan_cache_max_entries(),
1249 backend.runtime().stream_slot_count(),
1250 ))
1251 })?;
1252 Ok(Arc::clone(&guard))
1253}
1254
1255fn lock_cutensor_plan_cache(
1256 cache: &CutensorPlanCacheState,
1257) -> crate::Result<std::sync::MutexGuard<'_, CutensorContractionCacheState>> {
1258 cache
1259 .lock()
1260 .map_err(|_| Error::runtime_state("cutensor_plan_cache", "plan cache lock poisoned"))
1261}
1262
1263pub(super) fn cutensor_plan_cache_stats(backend: &CudaBackend) -> crate::Result<CacheStats> {
1264 let Some(plan_cache) = backend
1265 .cuda_extension_cache()
1266 .get_cloned::<CutensorPlanCacheState>()?
1267 else {
1268 return Ok(CacheStats::empty());
1269 };
1270 let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1271 Ok(plan_cache.plans.stats())
1272}
1273
1274pub(super) fn cutensor_workspace_stats(
1276 backend: &CudaBackend,
1277) -> crate::Result<CutensorWorkspaceStats> {
1278 let Some(plan_cache) = backend
1279 .cuda_extension_cache()
1280 .get_cloned::<CutensorPlanCacheState>()?
1281 else {
1282 return Ok(CutensorWorkspaceStats::default());
1283 };
1284 let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1285 Ok(plan_cache.workspace_stats())
1286}
1287
1288pub(super) fn set_cutensor_workspace_max_retained_bytes(
1294 backend: &CudaBackend,
1295 limit: u64,
1296) -> crate::Result<()> {
1297 let Some(plan_cache) = backend
1298 .cuda_extension_cache()
1299 .get_cloned::<CutensorPlanCacheState>()?
1300 else {
1301 return Ok(());
1302 };
1303 let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1304 if plan_cache.workspace_bytes() > limit {
1305 plan_cache.release_workspaces();
1306 }
1307 Ok(())
1308}
1309
1310pub(crate) fn cutensor_workspace_retirement_stats(
1312 backend: &CudaBackend,
1313) -> crate::Result<super::workspace_retirement::WorkspaceRetirementStats> {
1314 Ok(backend
1315 .runtime()
1316 .workspace_retirements()
1317 .lock()
1318 .map_err(|_| {
1319 crate::Error::runtime_state(
1320 "cutensor_workspace_retirement",
1321 "retirement queue lock poisoned",
1322 )
1323 })?
1324 .stats())
1325}
1326
1327pub(super) fn cutensor_plan_cache_max_entries(
1328 backend: &CudaBackend,
1329) -> crate::Result<NonZeroUsize> {
1330 let Some(plan_cache) = backend
1331 .cuda_extension_cache()
1332 .get_cloned::<CutensorPlanCacheState>()?
1333 else {
1334 return Ok(default_cutensor_plan_cache_max_entries());
1335 };
1336 let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1337 Ok(plan_cache.plans.max_entries())
1338}
1339
1340pub(super) fn set_cutensor_plan_cache_max_entries(
1341 backend: &CudaBackend,
1342 max_entries: NonZeroUsize,
1343) -> crate::Result<()> {
1344 let plan_cache = get_or_init_cutensor_plan_cache(backend)?;
1345 let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1346 plan_cache.plans.set_max_entries(max_entries);
1347 let retained_bytes = plan_cache.plans.retained_bytes();
1348 backend
1349 .cuda_extension_cache()
1350 .update_retained_bytes::<CutensorPlanCacheState>(retained_bytes)
1351}
1352
1353fn cached_cutensor_contraction<T, R>(
1354 backend: &CudaBackend,
1355 spec: &CutensorContractionSpec<'_>,
1356 cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
1357 execute: impl FnOnce(&CutensorHandle, &Plan, &Workspace) -> crate::Result<R>,
1358) -> crate::Result<R>
1359where
1360 T: CutensorScalar,
1361{
1362 let cutensor = backend.cutensor_handle()?;
1363 let hash = spec_hash::<T>(spec);
1364 let plan_cache = get_or_init_cutensor_plan_cache(backend)?;
1365 let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1366 let entries_changed = plan_cache.plans.ensure(
1367 hash,
1368 |key| key_matches_spec::<T>(key, spec),
1369 || {
1370 let cached = CachedCutensorContraction::new::<T>(cutensor, spec)?;
1371 let key = CutensorContractionKey::from_spec::<T>(spec);
1372 let retained_bytes = key.retained_bytes().saturating_add(cached.retained_bytes());
1373 Ok((key, cached, retained_bytes))
1374 },
1375 )?;
1376 if entries_changed {
1377 let retained_bytes = plan_cache.plans.retained_bytes();
1383 backend
1384 .cuda_extension_cache()
1385 .update_retained_bytes::<CutensorPlanCacheState>(retained_bytes)?;
1386 }
1387 let state = &mut *plan_cache;
1388 let slot = backend.runtime().stream_slot();
1389 let required = state
1390 .plans
1391 .get(hash, |key| key_matches_spec::<T>(key, spec))
1392 .map(|cached| cached.workspace_size)
1393 .ok_or_else(|| {
1394 Error::runtime_state(
1395 "cutensor_plan_cache",
1396 "cached cuTENSOR contraction was evicted before use",
1397 )
1398 })?;
1399 let current_capacity = state.workspaces[slot]
1400 .as_ref()
1401 .map_or(0, |workspace| workspace.size);
1402 let decision = plan_workspace(
1403 required,
1404 current_capacity,
1405 retained_workspace_bytes(&state.workspaces),
1406 backend.cutensor_workspace_limit(),
1407 );
1408 match decision {
1409 WorkspacePlan::Reuse => {}
1410 WorkspacePlan::Retain(capacity) => {
1411 let replacement = alloc_workspace(backend.runtime(), capacity)?;
1414 state.workspaces[slot] = Some(replacement);
1415 }
1416 WorkspacePlan::Temporary(capacity) => {
1417 backend.note_cutensor_temporary_workspace();
1422 let temporary = alloc_workspace(backend.runtime(), capacity)?;
1423 let cached = state
1424 .plans
1425 .get(hash, |key| key_matches_spec::<T>(key, spec))
1426 .ok_or_else(|| {
1427 Error::runtime_state(
1428 "cutensor_plan_cache",
1429 "cached cuTENSOR contraction was evicted before use",
1430 )
1431 })?;
1432 let execute_result = execute(cutensor, &cached.plan, &temporary);
1433 return backend.runtime().finish_vendor_enqueue(
1434 OP,
1435 cross_stream_handles,
1436 execute_result,
1437 );
1438 }
1439 }
1440 let cached = state
1441 .plans
1442 .get(hash, |key| key_matches_spec::<T>(key, spec))
1443 .ok_or_else(|| {
1444 Error::runtime_state(
1445 "cutensor_plan_cache",
1446 "cached cuTENSOR contraction was evicted before use",
1447 )
1448 })?;
1449 let empty = Workspace::none();
1452 let workspace = state.workspaces[slot].as_ref().unwrap_or(&empty);
1453 debug_assert!(workspace.size >= cached.workspace_size);
1454 let execute_result = execute(cutensor, &cached.plan, workspace);
1457 backend
1458 .runtime()
1459 .finish_vendor_enqueue(OP, cross_stream_handles, execute_result)
1460}
1461
1462pub(super) fn shared_workspace_capacity(requested: u64) -> Option<u64> {
1466 const MIN_CAPACITY: u64 = 1 << 20;
1467 if requested == 0 {
1468 return Some(0);
1469 }
1470 requested.max(MIN_CAPACITY).checked_next_power_of_two()
1471}
1472
1473fn validate_descriptor_alignment(
1474 actual_alignment: u32,
1475 alignment_requirement: u32,
1476 slot: &'static str,
1477) -> crate::Result<()> {
1478 if actual_alignment >= alignment_requirement {
1479 return Ok(());
1480 }
1481 Err(Error::invalid_argument(
1482 OP,
1483 "alignment",
1484 format!(
1485 "{slot} device pointer alignment {actual_alignment} is smaller than the cuTENSOR \
1486 descriptor requirement {alignment_requirement}"
1487 ),
1488 ))
1489}
1490
1491fn raw_stream(rt: &CudaRuntime) -> crate::Result<CutensorCudaStream> {
1492 Ok(rt.raw_cuda_stream()? as usize as CutensorCudaStream)
1493}
1494
1495fn alloc_workspace(rt: &CudaRuntime, workspace_size: u64) -> crate::Result<Workspace> {
1496 if workspace_size == 0 {
1497 return Ok(Workspace::none());
1498 }
1499 rt.workspace_retirements()
1503 .lock()
1504 .unwrap_or_else(|error| error.into_inner())
1505 .drain(rt.state());
1506 let workspace_len =
1507 usize::try_from(workspace_size).map_err(|_| workspace_size_overflow(OP, workspace_size))?;
1508 let handle = rt.client().empty(workspace_len);
1509 let resource = rt
1510 .client()
1511 .get_resource(handle.clone())
1512 .map_err(|err| crate::Error::backend_source(OP, err))?;
1513 let stream = rt.raw_cuda_stream()?;
1514 Ok(Workspace {
1515 _handle: Some(handle),
1516 ptr: cuda_device_ptr_from_addr(resource.resource().ptr, OP)?,
1517 size: workspace_size,
1518 runtime: Some(rt.clone()),
1519 stream,
1520 })
1521}
1522
1523pub(super) fn typed_device_ptr<T: TensorScalar + 'static>(
1524 rt: &CudaRuntime,
1525 tensor: &TypedTensor<T>,
1526 op: &'static str,
1527) -> crate::Result<*mut c_void> {
1528 ensure_resident_on_runtime(rt, tensor, op)?;
1529 let prepared = prepared_tensor_access(tensor, op)?;
1530 let buffer = cubecl_buffer(tensor, op)?;
1531 let addr = memoized_device_addr(rt, buffer, prepared, op)?;
1532 cuda_device_ptr_from_addr(addr, op)
1534}
1535
1536pub(super) fn write_device_ptr<T: TensorScalar + 'static>(
1555 rt: &CudaRuntime,
1556 tensor: &TypedTensor<T>,
1557 op: &'static str,
1558) -> crate::Result<*mut c_void> {
1559 ensure_resident_on_runtime(rt, tensor, op)?;
1560 let prepared = prepared_tensor_access(tensor, op)?;
1561 let resource = rt
1562 .client()
1563 .get_resource(prepared.into_handle())
1564 .map_err(|err| crate::Error::backend_source(op, err))?;
1565 let addr = resource.resource().ptr;
1566 cubecl_buffer(tensor, op)?.memoize_device_addr(addr);
1568 cuda_device_ptr_from_addr(addr, op)
1570}
1571
1572pub(super) fn memoized_device_addr(
1589 rt: &CudaRuntime,
1590 buffer: &CubeclBuffer,
1591 prepared: CubeclPreparedAccess,
1592 op: &'static str,
1593) -> crate::Result<u64> {
1594 let same_stream = StreamId::current() == buffer.handle().stream;
1595 if same_stream {
1596 if let Some(addr) = buffer.cached_device_addr() {
1597 return Ok(addr);
1598 }
1599 }
1600 let resource = rt
1601 .client()
1602 .get_resource(prepared.into_handle())
1603 .map_err(|err| crate::Error::backend_source(op, err))?;
1604 let addr = resource.resource().ptr;
1605 buffer.memoize_device_addr(addr);
1607 Ok(addr)
1608}
1609
1610fn build_layout(
1611 lhs_shape: &[usize],
1612 rhs_shape: &[usize],
1613 config: &DotGeneralConfig,
1614) -> crate::Result<DotGeneralLayout> {
1615 let lhs_free = free_axes(
1616 lhs_shape.len(),
1617 &config.lhs_contracting_dims,
1618 &config.lhs_batch_dims,
1619 );
1620 let rhs_free = free_axes(
1621 rhs_shape.len(),
1622 &config.rhs_contracting_dims,
1623 &config.rhs_batch_dims,
1624 );
1625
1626 let mut lhs_modes = vec![-1i32; lhs_shape.len()];
1627 let mut rhs_modes = vec![-1i32; rhs_shape.len()];
1628 let mut output_modes =
1629 Vec::with_capacity(lhs_free.len() + rhs_free.len() + config.lhs_batch_dims.len());
1630 let mut output_shape = Vec::with_capacity(output_modes.capacity());
1631 let mut batch_modes = Vec::with_capacity(config.lhs_batch_dims.len());
1632 let mut batch_shape = Vec::with_capacity(config.lhs_batch_dims.len());
1633 let mut next_mode = 0i32;
1634 let mut contracting_elements = 1usize;
1635
1636 for (&lhs_axis, &rhs_axis) in config
1637 .lhs_contracting_dims
1638 .iter()
1639 .zip(&config.rhs_contracting_dims)
1640 {
1641 let mode = next_mode;
1642 next_mode += 1;
1643 lhs_modes[lhs_axis] = mode;
1644 rhs_modes[rhs_axis] = mode;
1645 contracting_elements = contracting_elements
1646 .checked_mul(lhs_shape[lhs_axis])
1647 .ok_or_else(|| {
1648 Error::invalid_argument(
1649 OP,
1650 "shape",
1651 format!(
1652 "contracting dimension product overflows usize for lhs shape {lhs_shape:?}"
1653 ),
1654 )
1655 })?;
1656 }
1657
1658 for (&lhs_axis, &rhs_axis) in config.lhs_batch_dims.iter().zip(&config.rhs_batch_dims) {
1659 let mode = next_mode;
1660 next_mode += 1;
1661 lhs_modes[lhs_axis] = mode;
1662 rhs_modes[rhs_axis] = mode;
1663 batch_modes.push(mode);
1664 batch_shape.push(lhs_shape[lhs_axis]);
1665 }
1666
1667 for &lhs_axis in &lhs_free {
1668 let mode = next_mode;
1669 next_mode += 1;
1670 lhs_modes[lhs_axis] = mode;
1671 output_modes.push(mode);
1672 output_shape.push(lhs_shape[lhs_axis]);
1673 }
1674
1675 for &rhs_axis in &rhs_free {
1676 let mode = next_mode;
1677 next_mode += 1;
1678 rhs_modes[rhs_axis] = mode;
1679 output_modes.push(mode);
1680 output_shape.push(rhs_shape[rhs_axis]);
1681 }
1682
1683 output_modes.extend_from_slice(&batch_modes);
1684 output_shape.extend_from_slice(&batch_shape);
1685
1686 let lhs_extents = dims_to_i64(lhs_shape)?;
1687 let rhs_extents = dims_to_i64(rhs_shape)?;
1688 let output_extents = dims_to_i64(&output_shape)?;
1689 let lhs_strides = strides_to_i64(&col_major_strides(lhs_shape)?)?;
1690 let rhs_strides = strides_to_i64(&col_major_strides(rhs_shape)?)?;
1691 let output_strides = strides_to_i64(&col_major_strides(&output_shape)?)?;
1692
1693 Ok(DotGeneralLayout {
1694 lhs_modes,
1695 rhs_modes,
1696 output_modes,
1697 output_shape,
1698 lhs_extents,
1699 rhs_extents,
1700 output_extents,
1701 lhs_strides,
1702 rhs_strides,
1703 output_strides,
1704 contracting_elements,
1705 })
1706}
1707
1708fn dims_to_i64(dims: &[usize]) -> crate::Result<Vec<i64>> {
1709 dims.iter()
1710 .map(|&dim| {
1711 i64::try_from(dim).map_err(|_| {
1712 Error::invalid_argument(
1713 OP,
1714 "shape",
1715 format!("extent {dim} exceeds cuTENSOR i64 limit"),
1716 )
1717 })
1718 })
1719 .collect()
1720}
1721
1722fn strides_to_i64(strides: &[isize]) -> crate::Result<Vec<i64>> {
1723 strides
1724 .iter()
1725 .map(|&stride| {
1726 i64::try_from(stride).map_err(|_| {
1727 Error::invalid_argument(
1728 OP,
1729 "stride",
1730 format!("stride {stride} exceeds cuTENSOR i64 limit"),
1731 )
1732 })
1733 })
1734 .collect()
1735}
1736
1737fn free_axes(rank: usize, contracting: &[usize], batch: &[usize]) -> Vec<usize> {
1738 (0..rank)
1739 .filter(|axis| !contracting.contains(axis) && !batch.contains(axis))
1740 .collect()
1741}
1742
1743fn validate_axis_list(
1744 op: &'static str,
1745 role: &'static str,
1746 axes: &[usize],
1747 rank: usize,
1748) -> crate::Result<()> {
1749 let mut seen = vec![false; rank];
1750 for &axis in axes {
1751 if axis >= rank {
1752 return Err(Error::axis_out_of_bounds(op, axis, rank));
1753 }
1754 if seen[axis] {
1755 return Err(Error::duplicate_axis(op, axis, role));
1756 }
1757 seen[axis] = true;
1758 }
1759 Ok(())
1760}
1761
1762fn validate_role_disjoint(
1763 op: &'static str,
1764 first_role: &'static str,
1765 first_axes: &[usize],
1766 second_role: &'static str,
1767 second_axes: &[usize],
1768) -> crate::Result<()> {
1769 for &axis in first_axes {
1770 if second_axes.contains(&axis) {
1771 return Err(Error::validation(
1772 op,
1773 tenferro_tensor::ValidationError::AxisRoleConflict {
1774 axis,
1775 first_role,
1776 second_role,
1777 },
1778 ));
1779 }
1780 }
1781 Ok(())
1782}
1783
1784fn validate_dot_general(
1785 lhs_shape: &[usize],
1786 rhs_shape: &[usize],
1787 config: &DotGeneralConfig,
1788) -> crate::Result<()> {
1789 if config.lhs_contracting_dims.len() != config.rhs_contracting_dims.len() {
1790 return Err(Error::invalid_argument(
1791 OP,
1792 "contracting_dims",
1793 "lhs/rhs contracting dim counts differ",
1794 ));
1795 }
1796 if config.lhs_batch_dims.len() != config.rhs_batch_dims.len() {
1797 return Err(Error::invalid_argument(
1798 OP,
1799 "batch_dims",
1800 "lhs/rhs batch dim counts differ",
1801 ));
1802 }
1803
1804 let lhs_rank = lhs_shape.len();
1805 let rhs_rank = rhs_shape.len();
1806
1807 validate_axis_list(
1808 OP,
1809 "lhs_contracting",
1810 &config.lhs_contracting_dims,
1811 lhs_rank,
1812 )?;
1813 validate_axis_list(
1814 OP,
1815 "rhs_contracting",
1816 &config.rhs_contracting_dims,
1817 rhs_rank,
1818 )?;
1819 validate_axis_list(OP, "lhs_batch", &config.lhs_batch_dims, lhs_rank)?;
1820 validate_axis_list(OP, "rhs_batch", &config.rhs_batch_dims, rhs_rank)?;
1821 validate_role_disjoint(
1822 OP,
1823 "lhs_contracting",
1824 &config.lhs_contracting_dims,
1825 "lhs_batch",
1826 &config.lhs_batch_dims,
1827 )?;
1828 validate_role_disjoint(
1829 OP,
1830 "rhs_contracting",
1831 &config.rhs_contracting_dims,
1832 "rhs_batch",
1833 &config.rhs_batch_dims,
1834 )?;
1835
1836 for (&lhs_axis, &rhs_axis) in config
1837 .lhs_contracting_dims
1838 .iter()
1839 .zip(&config.rhs_contracting_dims)
1840 {
1841 if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
1842 return Err(Error::validation(
1843 OP,
1844 tenferro_tensor::ShapeMismatch::ContractedDimensions {
1845 lhs_axis,
1846 lhs_size: lhs_shape[lhs_axis],
1847 rhs_axis,
1848 rhs_size: rhs_shape[rhs_axis],
1849 }
1850 .into(),
1851 ));
1852 }
1853 }
1854
1855 for (&lhs_axis, &rhs_axis) in config.lhs_batch_dims.iter().zip(&config.rhs_batch_dims) {
1856 if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
1857 return Err(Error::shape_mismatch(
1858 OP,
1859 lhs_shape.to_vec(),
1860 rhs_shape.to_vec(),
1861 ));
1862 }
1863 }
1864
1865 Ok(())
1866}