1use super::{DType, Tensor};
6use crate::errors::{Result, TrustformersError};
7use scirs2_core::ndarray::{ArrayD, IxDyn};
8use std::collections::HashMap;
9use std::sync::atomic::AtomicU64;
10use std::sync::{Arc, RwLock};
11
12#[allow(dead_code)] static TENSOR_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
15
16lazy_static::lazy_static! {
17 static ref GRADIENT_REGISTRY: Arc<RwLock<HashMap<u64, Tensor>>> = Arc::new(RwLock::new(HashMap::new()));
19}
20
21thread_local! {
22 static GRADIENT_MODE: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
24}
25
26pub fn enable_grad() {
28 GRADIENT_MODE.with(|mode| mode.set(true));
29}
30
31pub fn disable_grad() {
33 GRADIENT_MODE.with(|mode| mode.set(false));
34}
35
36pub fn is_grad_enabled() -> bool {
38 GRADIENT_MODE.with(|mode| mode.get())
39}
40
41pub fn clear_gradients() {
43 if let Ok(mut registry) = GRADIENT_REGISTRY.write() {
44 registry.clear();
45 }
46}
47
48impl Tensor {
49 fn tensor_id(&self) -> u64 {
51 use std::collections::hash_map::DefaultHasher;
53 use std::hash::{Hash, Hasher};
54
55 let mut hasher = DefaultHasher::new();
56 self.shape().hash(&mut hasher);
57
58 match self {
59 Tensor::F32(arr) => {
60 arr.as_ptr().hash(&mut hasher);
61 arr.len().hash(&mut hasher);
62 },
63 Tensor::F64(arr) => {
64 arr.as_ptr().hash(&mut hasher);
65 arr.len().hash(&mut hasher);
66 },
67 Tensor::I64(arr) => {
68 arr.as_ptr().hash(&mut hasher);
69 arr.len().hash(&mut hasher);
70 },
71 #[cfg(all(target_os = "macos", feature = "metal"))]
72 Tensor::Metal(data) => {
73 data.buffer_id.hash(&mut hasher);
75 self.len().hash(&mut hasher);
76 },
77 #[cfg(feature = "cuda")]
78 Tensor::CUDA(data) => {
79 data.buffer_id().hash(&mut hasher);
81 self.len().hash(&mut hasher);
82 },
83 _ => {
84 self.len().hash(&mut hasher);
86 },
87 }
88
89 hasher.finish()
90 }
91 pub fn shape(&self) -> Vec<usize> {
97 match self {
98 Tensor::F32(a) => a.shape().to_vec(),
99 Tensor::F64(a) => a.shape().to_vec(),
100 Tensor::F16(a) => a.shape().to_vec(),
101 Tensor::BF16(a) => a.shape().to_vec(),
102 Tensor::I64(a) => a.shape().to_vec(),
103 Tensor::C32(a) => a.shape().to_vec(),
104 Tensor::C64(a) => a.shape().to_vec(),
105 Tensor::CF16(a) => a.shape().to_vec(),
106 Tensor::CBF16(a) => a.shape().to_vec(),
107 Tensor::Sparse(s) => s.shape().to_vec(),
108 #[cfg(feature = "candle")]
109 Tensor::Candle(t) => t.shape().dims().to_vec(),
110 #[cfg(all(target_os = "macos", feature = "metal"))]
111 Tensor::Metal(data) => data.shape.clone(),
112 #[cfg(feature = "cuda")]
113 Tensor::CUDA(data) => data.shape.clone(),
114 }
115 }
116
117 pub fn len(&self) -> usize {
123 match self {
124 Tensor::F32(a) => a.len(),
125 Tensor::F64(a) => a.len(),
126 Tensor::F16(a) => a.len(),
127 Tensor::BF16(a) => a.len(),
128 Tensor::I64(a) => a.len(),
129 Tensor::C32(a) => a.len(),
130 Tensor::C64(a) => a.len(),
131 Tensor::CF16(a) => a.len(),
132 Tensor::CBF16(a) => a.len(),
133 Tensor::Sparse(s) => s.nnz(), #[cfg(feature = "candle")]
135 Tensor::Candle(t) => t.elem_count(),
136 #[cfg(all(target_os = "macos", feature = "metal"))]
137 Tensor::Metal(data) => data.shape.iter().product(),
138 #[cfg(feature = "cuda")]
139 Tensor::CUDA(data) => data.shape.iter().product(),
140 }
141 }
142
143 pub fn is_empty(&self) -> bool {
149 self.len() == 0
150 }
151
152 pub fn ndim(&self) -> usize {
158 self.shape().len()
159 }
160
161 pub fn size_bytes(&self) -> usize {
167 match self {
168 Tensor::F32(a) => a.len() * std::mem::size_of::<f32>(),
169 Tensor::F64(a) => a.len() * std::mem::size_of::<f64>(),
170 Tensor::F16(a) => a.len() * std::mem::size_of::<half::f16>(),
171 Tensor::BF16(a) => a.len() * std::mem::size_of::<half::bf16>(),
172 Tensor::I64(a) => a.len() * std::mem::size_of::<i64>(),
173 Tensor::C32(a) => a.len() * std::mem::size_of::<scirs2_core::Complex32>(),
174 Tensor::C64(a) => a.len() * std::mem::size_of::<scirs2_core::Complex64>(),
175 Tensor::CF16(a) => a.len() * std::mem::size_of::<scirs2_core::Complex<half::f16>>(),
176 Tensor::CBF16(a) => a.len() * std::mem::size_of::<scirs2_core::Complex<half::bf16>>(),
177 Tensor::Sparse(s) => s.nnz() * std::mem::size_of::<f32>(), #[cfg(feature = "candle")]
179 Tensor::Candle(t) => t.elem_count() * std::mem::size_of::<f32>(), #[cfg(all(target_os = "macos", feature = "metal"))]
181 Tensor::Metal(data) => {
182 let num_elements: usize = data.shape.iter().product();
183 num_elements * data.dtype.size_in_bytes()
184 },
185 #[cfg(feature = "cuda")]
186 Tensor::CUDA(data) => {
187 let num_elements: usize = data.shape.iter().product();
188 num_elements * data.dtype.size_in_bytes()
189 },
190 }
191 }
192
193 pub fn to_device(&self, device: &str) -> Result<Tensor> {
203 let device_lower = device.to_lowercase();
205
206 let (device_type, device_index) = if device_lower.contains(':') {
208 let parts: Vec<&str> = device_lower.split(':').collect();
209 if parts.len() != 2 {
210 return Err(TrustformersError::tensor_op_error(
211 &format!("Invalid device format '{}'. Expected format: 'device_type' or 'device_type:index'", device),
212 "to_device"
213 ));
214 }
215
216 let index = parts[1].parse::<usize>().map_err(|_| {
217 TrustformersError::tensor_op_error(
218 &format!(
219 "Invalid device index '{}'. Expected a non-negative integer",
220 parts[1]
221 ),
222 "to_device",
223 )
224 })?;
225
226 (parts[0], Some(index))
227 } else {
228 (device_lower.as_str(), None)
229 };
230
231 match device_type {
233 "cpu" => {
234 if let Some(index) = device_index {
236 if index > 0 {
237 return Err(TrustformersError::tensor_op_error(
238 &format!("CPU device index {} not supported. CPU only supports index 0 or no index", index),
239 "to_device"
240 ));
241 }
242 }
243 Ok(self.clone())
245 },
246 "cuda" => {
247 if let Some(index) = device_index {
249 Err(TrustformersError::tensor_op_error(
250 &format!("CUDA device cuda:{} not available. This build doesn't support CUDA. Consider using CPU instead with device='cpu'", index),
251 "to_device"
252 ))
253 } else {
254 Err(TrustformersError::tensor_op_error(
255 "CUDA devices not available. This build doesn't support CUDA. Consider using CPU instead with device='cpu'",
256 "to_device"
257 ))
258 }
259 },
260 "mps" => {
261 Err(TrustformersError::tensor_op_error(
263 "MPS device not available. This build doesn't support Metal Performance Shaders. Consider using CPU instead with device='cpu'",
264 "to_device"
265 ))
266 },
267 "tpu" => {
268 Err(TrustformersError::tensor_op_error(
270 "TPU devices not available. This build doesn't support TPU. Consider using CPU instead with device='cpu'",
271 "to_device"
272 ))
273 },
274 "xpu" | "intel" => {
275 Err(TrustformersError::tensor_op_error(
277 "Intel XPU devices not available. This build doesn't support Intel XPU. Consider using CPU instead with device='cpu'",
278 "to_device"
279 ))
280 },
281 "npu" => {
282 Err(TrustformersError::tensor_op_error(
284 "NPU devices not available. This build doesn't support NPU. Consider using CPU instead with device='cpu'",
285 "to_device"
286 ))
287 },
288 _ => {
289 Err(TrustformersError::tensor_op_error(
290 &format!("Unknown device type '{}'. Supported device types: cpu, cuda, mps, tpu, xpu, npu. For this build, only 'cpu' is supported", device_type),
291 "to_device"
292 ))
293 },
294 }
295 }
296
297 pub fn to_device_enum(&self, device: &crate::device::Device) -> Result<Tensor> {
328 match (self, device) {
329 #[cfg(all(target_os = "macos", feature = "metal"))]
331 (Tensor::F32(arr), crate::device::Device::Metal(_)) => {
332 use crate::gpu_ops::metal::get_metal_backend;
333 let backend = get_metal_backend()?;
334 let data_vec: Vec<f32> = arr.iter().copied().collect();
335
336 #[cfg(debug_assertions)]
337 {
338 eprintln!(
340 "🔍 to_device_enum(F32→Metal): data_vec.len()={}",
341 data_vec.len()
342 );
343 if !data_vec.is_empty() {
344 eprintln!(
345 "🔍 to_device_enum: first 10 values: {:?}",
346 &data_vec[..10.min(data_vec.len())]
347 );
348 eprintln!(
349 "🔍 to_device_enum: stats - min={:.4}, max={:.4}, mean={:.4}",
350 data_vec.iter().fold(f32::INFINITY, |a, &b| a.min(b)),
351 data_vec.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b)),
352 data_vec.iter().sum::<f32>() / data_vec.len() as f32
353 );
354 }
355 }
356
357 let buffer_id = backend.create_persistent_buffer(&data_vec)?;
358
359 #[cfg(debug_assertions)]
360 {
361 eprintln!("🔍 to_device_enum: Created buffer_id={:?}", buffer_id);
362
363 let verify_data = backend.download_buffer_to_vec(&buffer_id)?;
365 eprintln!(
366 "🔍 to_device_enum: Verification download - len={}, first 10: {:?}",
367 verify_data.len(),
368 &verify_data[..10.min(verify_data.len())]
369 );
370 }
371
372 Ok(Tensor::Metal(super::MetalTensorData {
373 buffer_id,
374 shape: arr.shape().to_vec(),
375 dtype: DType::F32,
376 }))
377 },
378
379 #[cfg(all(target_os = "macos", feature = "metal"))]
381 (Tensor::F64(arr), crate::device::Device::Metal(_)) => {
382 use crate::gpu_ops::metal::get_metal_backend;
383 let backend = get_metal_backend()?;
384 let data_vec: Vec<f32> = arr.iter().map(|&x| x as f32).collect();
385 let buffer_id = backend.create_persistent_buffer(&data_vec)?;
386 Ok(Tensor::Metal(super::MetalTensorData {
387 buffer_id,
388 shape: arr.shape().to_vec(),
389 dtype: DType::F32,
390 }))
391 },
392
393 #[cfg(all(target_os = "macos", feature = "metal"))]
395 (Tensor::Metal(metal_data), crate::device::Device::CPU) => {
396 use crate::gpu_ops::metal::get_metal_backend;
397 let backend = get_metal_backend()?;
398 let buffer = backend.get_persistent_buffer(&metal_data.buffer_id)?;
399
400 let size: usize = metal_data.shape.iter().product();
402
403 match metal_data.dtype {
405 DType::F32 => {
406 let ptr = buffer.contents() as *const f32;
407 let data_vec = unsafe { std::slice::from_raw_parts(ptr, size) }.to_vec();
408
409 use scirs2_core::ndarray::ArrayD;
411 let arr = ArrayD::from_shape_vec(
412 scirs2_core::ndarray::IxDyn(&metal_data.shape),
413 data_vec,
414 )
415 .map_err(|e| {
416 TrustformersError::tensor_op_error(
417 &format!("Failed to create array from shape: {}", e),
418 "to_device_enum",
419 )
420 })?;
421 Ok(Tensor::F32(arr))
422 },
423 _ => Err(TrustformersError::tensor_op_error(
424 &format!("Unsupported Metal tensor dtype: {:?}", metal_data.dtype),
425 "to_device_enum",
426 )),
427 }
428 },
429
430 #[cfg(all(target_os = "macos", feature = "metal"))]
432 (Tensor::Metal(metal_data), crate::device::Device::Metal(_)) => {
433 Ok(Tensor::Metal(metal_data.clone()))
436 },
437
438 (Tensor::F32(_), crate::device::Device::CPU) => Ok(self.clone()),
440 (Tensor::F64(_), crate::device::Device::CPU) => Ok(self.clone()),
441 (Tensor::F16(_), crate::device::Device::CPU) => Ok(self.clone()),
442 (Tensor::BF16(_), crate::device::Device::CPU) => Ok(self.clone()),
443 (Tensor::I64(_), crate::device::Device::CPU) => Ok(self.clone()),
444 (Tensor::C32(_), crate::device::Device::CPU) => Ok(self.clone()),
445 (Tensor::C64(_), crate::device::Device::CPU) => Ok(self.clone()),
446 (Tensor::CF16(_), crate::device::Device::CPU) => Ok(self.clone()),
447 (Tensor::CBF16(_), crate::device::Device::CPU) => Ok(self.clone()),
448 (Tensor::Sparse(_), crate::device::Device::CPU) => Ok(self.clone()),
449
450 #[cfg(not(feature = "metal"))]
452 (_, crate::device::Device::Metal(_)) => Err(TrustformersError::hardware_error(
453 "Metal not available. Compile with --features metal",
454 "to_device_enum",
455 )),
456
457 #[cfg(feature = "cuda")]
459 #[allow(unused_variables)]
460 (Tensor::F32(arr), crate::device::Device::CUDA(device_id)) => {
461 #[cfg(any(target_os = "linux", target_os = "windows"))]
462 {
463 use crate::gpu_ops::cuda::get_cuda_backend;
464 let backend = get_cuda_backend(*device_id)?;
465 let data_vec: Vec<f32> = arr.iter().copied().collect();
466 let buffer_id = backend.create_persistent_buffer(&data_vec)?;
467 Ok(Tensor::CUDA(super::CudaTensorData::new(
468 buffer_id,
469 *device_id,
470 arr.shape().to_vec(),
471 DType::F32,
472 )))
473 }
474 #[cfg(not(any(target_os = "linux", target_os = "windows")))]
475 {
476 Err(TrustformersError::hardware_error(
477 "CUDA is only supported on Linux and Windows",
478 "to_device_enum",
479 ))
480 }
481 },
482
483 #[cfg(feature = "cuda")]
485 #[allow(unused_variables)]
486 (Tensor::F64(arr), crate::device::Device::CUDA(device_id)) => {
487 #[cfg(any(target_os = "linux", target_os = "windows"))]
488 {
489 use crate::gpu_ops::cuda::get_cuda_backend;
490 let backend = get_cuda_backend(*device_id)?;
491 let data_vec: Vec<f32> = arr.iter().map(|&x| x as f32).collect();
492 let buffer_id = backend.create_persistent_buffer(&data_vec)?;
493 Ok(Tensor::CUDA(super::CudaTensorData::new(
494 buffer_id,
495 *device_id,
496 arr.shape().to_vec(),
497 DType::F32,
498 )))
499 }
500 #[cfg(not(any(target_os = "linux", target_os = "windows")))]
501 {
502 Err(TrustformersError::hardware_error(
503 "CUDA is only supported on Linux and Windows",
504 "to_device_enum",
505 ))
506 }
507 },
508
509 #[cfg(feature = "cuda")]
511 #[allow(unused_variables)]
512 (Tensor::CUDA(cuda_data), crate::device::Device::CPU) => {
513 #[cfg(any(target_os = "linux", target_os = "windows"))]
514 {
515 use crate::gpu_ops::cuda::get_cuda_backend;
516 let backend = get_cuda_backend(cuda_data.device_id())?;
519
520 match cuda_data.dtype {
522 DType::F32 => {
523 let data_vec = backend.download_buffer(&cuda_data.buffer_id())?;
525
526 use scirs2_core::ndarray::ArrayD;
528 let arr = ArrayD::from_shape_vec(
529 scirs2_core::ndarray::IxDyn(&cuda_data.shape),
530 data_vec,
531 )
532 .map_err(|e| {
533 TrustformersError::tensor_op_error(
534 &format!("Failed to create array from shape: {}", e),
535 "to_device_enum",
536 )
537 })?;
538 Ok(Tensor::F32(arr))
539 },
540 _ => Err(TrustformersError::tensor_op_error(
541 &format!("Unsupported CUDA tensor dtype: {:?}", cuda_data.dtype),
542 "to_device_enum",
543 )),
544 }
545 }
546 #[cfg(not(any(target_os = "linux", target_os = "windows")))]
547 {
548 Err(TrustformersError::hardware_error(
549 "CUDA is only supported on Linux and Windows",
550 "to_device_enum",
551 ))
552 }
553 },
554
555 #[cfg(feature = "cuda")]
557 (Tensor::CUDA(cuda_data), crate::device::Device::CUDA(target_device)) => {
558 if *target_device == cuda_data.device_id() {
559 Ok(Tensor::CUDA(cuda_data.clone()))
562 } else {
563 let host_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
566 host_tensor.to_device_enum(device)
567 }
568 },
569
570 #[cfg(not(feature = "cuda"))]
572 (_, crate::device::Device::CUDA(_)) => Err(TrustformersError::hardware_error(
573 "CUDA not available. Compile with --features cuda",
574 "to_device_enum",
575 )),
576
577 (_, crate::device::Device::ROCm(_)) => Err(TrustformersError::hardware_error(
579 "ROCm transfer not implemented yet",
580 "to_device_enum",
581 )),
582
583 (_, crate::device::Device::WebGPU) => Err(TrustformersError::hardware_error(
585 "WebGPU transfer not implemented yet",
586 "to_device_enum",
587 )),
588
589 #[allow(unreachable_patterns)]
591 _ => Err(TrustformersError::tensor_op_error(
592 &format!(
593 "Unsupported device transfer from {:?} to {:?}",
594 self.dtype(),
595 device
596 ),
597 "to_device_enum",
598 )),
599 }
600 }
601
602 pub fn grad(&self) -> Result<Tensor> {
626 if !is_grad_enabled() {
627 return Err(TrustformersError::tensor_op_error(
628 "Gradient tracking is not enabled. Use enable_grad() to enable gradient tracking.",
629 "grad",
630 ));
631 }
632
633 let tensor_id = self.tensor_id();
634
635 if let Ok(registry) = GRADIENT_REGISTRY.read() {
636 if let Some(grad_tensor) = registry.get(&tensor_id) {
637 Ok(grad_tensor.clone())
638 } else {
639 Err(TrustformersError::tensor_op_error(
640 "No gradient found for this tensor. Gradients are set during backward pass.",
641 "grad",
642 ))
643 }
644 } else {
645 Err(TrustformersError::tensor_op_error(
646 "Failed to access gradient registry.",
647 "grad",
648 ))
649 }
650 }
651
652 pub fn set_grad(&mut self, grad: Tensor) -> Result<()> {
677 if !is_grad_enabled() {
678 return Err(TrustformersError::tensor_op_error(
679 "Gradient tracking is not enabled. Use enable_grad() to enable gradient tracking.",
680 "set_grad",
681 ));
682 }
683
684 if self.shape() != grad.shape() {
686 return Err(TrustformersError::tensor_op_error(
687 &format!(
688 "Gradient shape {:?} doesn't match tensor shape {:?}",
689 grad.shape(),
690 self.shape()
691 ),
692 "set_grad",
693 ));
694 }
695
696 let tensor_id = self.tensor_id();
697
698 if let Ok(mut registry) = GRADIENT_REGISTRY.write() {
699 registry.insert(tensor_id, grad);
700 Ok(())
701 } else {
702 Err(TrustformersError::tensor_op_error(
703 "Failed to access gradient registry.",
704 "set_grad",
705 ))
706 }
707 }
708
709 pub fn data(&self) -> Result<Vec<f32>> {
715 match self {
716 Tensor::F32(a) => Ok(a.iter().cloned().collect()),
717 Tensor::F64(a) => Ok(a.iter().map(|&x| x as f32).collect()),
718 Tensor::I64(a) => Ok(a.iter().map(|&x| x as f32).collect()),
719 #[cfg(all(target_os = "macos", feature = "metal"))]
720 Tensor::Metal(_) => {
721 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
723 cpu_tensor.data()
724 },
725 #[cfg(feature = "cuda")]
726 Tensor::CUDA(_) => {
727 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
729 cpu_tensor.data()
730 },
731 _ => Err(TrustformersError::tensor_op_error(
732 "Unsupported tensor type for data conversion",
733 "data_conversion",
734 )),
735 }
736 }
737
738 pub fn softmax_entropy_normalized(&self) -> Result<f32> {
745 let values = self.to_vec_f32()?;
746 let n = values.len();
747 if n <= 1 {
748 return Ok(0.0);
749 }
750 let max = values.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
751 let exps: Vec<f32> = values.iter().map(|&v| (v - max).exp()).collect();
752 let sum: f32 = exps.iter().sum();
753 if sum <= 0.0 {
754 return Ok(0.0);
755 }
756 let entropy: f32 =
757 exps.iter().map(|&e| e / sum).filter(|&p| p > 1e-10).map(|p| -p * p.ln()).sum();
758 Ok((entropy / (n as f32).ln()).clamp(0.0, 1.0))
759 }
760
761 pub fn data_f32(&self) -> Result<Vec<f32>> {
767 self.data()
768 }
769
770 pub fn set_data_f32(&mut self, data: &[f32]) -> Result<()> {
780 match self {
781 Tensor::F32(a) => {
782 let shape = a.shape().to_vec();
783 let expected_len: usize = shape.iter().product();
784 if data.len() != expected_len {
785 return Err(TrustformersError::tensor_op_error(
786 &format!(
787 "Data length {} does not match tensor size {}",
788 data.len(),
789 expected_len
790 ),
791 "set_data_f32",
792 ));
793 }
794 *a = ArrayD::from_shape_vec(IxDyn(&shape), data.to_vec()).map_err(|e| {
795 TrustformersError::tensor_op_error(&e.to_string(), "set_data_f32")
796 })?;
797 Ok(())
798 },
799 _ => Err(TrustformersError::tensor_op_error(
800 "set_data_f32 only supported for F32 tensors",
801 "set_data_f32",
802 )),
803 }
804 }
805
806 pub fn data_mut(&mut self) -> Result<&mut [f32]> {
812 match self {
813 Tensor::F32(a) => a.as_slice_mut().ok_or_else(|| {
814 TrustformersError::tensor_op_error(
815 "Tensor data must be contiguous for mutable access",
816 "data_mut",
817 )
818 }),
819 _ => Err(TrustformersError::tensor_op_error(
820 "Mutable data access only supported for F32 tensors",
821 "data_mut",
822 )),
823 }
824 }
825
826 pub fn modify_data<F>(&mut self, f: F) -> Result<()>
836 where
837 F: FnOnce(&mut [f32]),
838 {
839 match self {
840 Tensor::F32(a) => {
841 if let Some(slice) = a.as_slice_mut() {
842 f(slice);
843 Ok(())
844 } else {
845 Err(TrustformersError::tensor_op_error(
846 "Cannot get mutable slice",
847 "modify_data",
848 ))
849 }
850 },
851 _ => Err(TrustformersError::tensor_op_error(
852 "Modify data only supported for F32 tensors",
853 "modify_data",
854 )),
855 }
856 }
857
858 pub fn device(&self) -> String {
864 match self {
865 Tensor::F32(_)
866 | Tensor::F64(_)
867 | Tensor::F16(_)
868 | Tensor::BF16(_)
869 | Tensor::I64(_)
870 | Tensor::C32(_)
871 | Tensor::C64(_)
872 | Tensor::CF16(_)
873 | Tensor::CBF16(_) => "cpu".to_string(),
874 Tensor::Sparse(_) => "cpu".to_string(),
875 #[cfg(feature = "candle")]
876 Tensor::Candle(t) => format!("{:?}", t.device()),
877 #[cfg(all(target_os = "macos", feature = "metal"))]
878 Tensor::Metal(_) => "metal".to_string(),
879 #[cfg(feature = "cuda")]
880 Tensor::CUDA(_) => "cuda".to_string(),
881 }
882 }
883
884 pub fn size(&self) -> usize {
890 self.shape().iter().product()
891 }
892
893 pub fn memory_usage(&self) -> usize {
899 match self {
900 Tensor::F32(a) => a.len() * std::mem::size_of::<f32>(),
901 Tensor::F64(a) => a.len() * std::mem::size_of::<f64>(),
902 Tensor::F16(a) => a.len() * std::mem::size_of::<half::f16>(),
903 Tensor::BF16(a) => a.len() * std::mem::size_of::<half::bf16>(),
904 Tensor::I64(a) => a.len() * std::mem::size_of::<i64>(),
905 Tensor::C32(a) => a.len() * std::mem::size_of::<scirs2_core::Complex32>(),
906 Tensor::C64(a) => a.len() * std::mem::size_of::<scirs2_core::Complex64>(),
907 Tensor::CF16(a) => a.len() * std::mem::size_of::<scirs2_core::Complex<half::f16>>(),
908 Tensor::CBF16(a) => a.len() * std::mem::size_of::<scirs2_core::Complex<half::bf16>>(),
909 Tensor::Sparse(s) => s.memory_usage(),
910 #[cfg(feature = "candle")]
911 Tensor::Candle(t) => t.elem_count() * 4, #[cfg(all(target_os = "macos", feature = "metal"))]
913 Tensor::Metal(m) => m.shape.iter().product::<usize>() * 4, #[cfg(feature = "cuda")]
915 Tensor::CUDA(c) => c.shape.iter().product::<usize>() * 4, }
917 }
918
919 pub fn dtype(&self) -> DType {
925 match self {
926 Tensor::F32(_) => DType::F32,
927 Tensor::F64(_) => DType::F64,
928 Tensor::F16(_) => DType::F16,
929 Tensor::BF16(_) => DType::BF16,
930 Tensor::I64(_) => DType::I64,
931 Tensor::C32(_) => DType::C32,
932 Tensor::C64(_) => DType::C64,
933 Tensor::CF16(_) => DType::CF16,
934 Tensor::CBF16(_) => DType::CBF16,
935 Tensor::Sparse(_) => DType::F32, #[cfg(feature = "candle")]
937 Tensor::Candle(_) => DType::F32, #[cfg(all(target_os = "macos", feature = "metal"))]
939 Tensor::Metal(data) => data.dtype,
940 #[cfg(feature = "cuda")]
941 Tensor::CUDA(data) => data.dtype,
942 }
943 }
944
945 pub fn get_dtype(&self) -> DType {
947 self.dtype()
948 }
949
950 pub fn get_float(&self, index: usize) -> Result<f32> {
960 match self {
961 Tensor::F32(a) => {
962 if index >= a.len() {
963 return Err(TrustformersError::tensor_op_error(
964 &format!(
965 "Index {} out of bounds for tensor of size {}",
966 index,
967 a.len()
968 ),
969 "get_float",
970 ));
971 }
972 Ok(a.iter().nth(index).copied().unwrap_or(0.0))
973 },
974 Tensor::F64(a) => {
975 if index >= a.len() {
976 return Err(TrustformersError::tensor_op_error(
977 &format!(
978 "Index {} out of bounds for tensor of size {}",
979 index,
980 a.len()
981 ),
982 "get_float",
983 ));
984 }
985 Ok(a.iter().nth(index).copied().unwrap_or(0.0) as f32)
986 },
987 #[cfg(all(target_os = "macos", feature = "metal"))]
988 Tensor::Metal(_) => {
989 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
991 cpu_tensor.get_float(index)
992 },
993 #[cfg(feature = "cuda")]
994 Tensor::CUDA(_) => {
995 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
997 cpu_tensor.get_float(index)
998 },
999 _ => Err(TrustformersError::tensor_op_error(
1000 "Get float not supported for this tensor type",
1001 "get_float",
1002 )),
1003 }
1004 }
1005
1006 pub fn item<T>(&self) -> Result<T>
1016 where
1017 T: num_traits::NumCast,
1018 {
1019 if self.len() != 1 {
1020 return Err(TrustformersError::tensor_op_error(
1021 &format!(
1022 "item() requires a single-element tensor, but got {} elements",
1023 self.len()
1024 ),
1025 "item",
1026 ));
1027 }
1028
1029 match self {
1030 Tensor::F32(a) => {
1031 let val = a.iter().next().copied().unwrap_or(0.0);
1032 T::from(val).ok_or_else(|| {
1033 TrustformersError::tensor_op_error(
1034 "Failed to convert f32 to target type",
1035 "item",
1036 )
1037 })
1038 },
1039 Tensor::F64(a) => {
1040 let val = a.iter().next().copied().unwrap_or(0.0);
1041 T::from(val).ok_or_else(|| {
1042 TrustformersError::tensor_op_error(
1043 "Failed to convert f64 to target type",
1044 "item",
1045 )
1046 })
1047 },
1048 Tensor::I64(a) => {
1049 let val = a.iter().next().copied().unwrap_or(0);
1050 T::from(val).ok_or_else(|| {
1051 TrustformersError::tensor_op_error(
1052 "Failed to convert i64 to target type",
1053 "item",
1054 )
1055 })
1056 },
1057 #[cfg(all(target_os = "macos", feature = "metal"))]
1058 Tensor::Metal(_) => {
1059 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
1061 cpu_tensor.item::<T>()
1062 },
1063 #[cfg(feature = "cuda")]
1064 Tensor::CUDA(_) => {
1065 let cpu_tensor = self.to_device_enum(&crate::device::Device::CPU)?;
1067 cpu_tensor.item::<T>()
1068 },
1069 _ => Err(TrustformersError::tensor_op_error(
1070 "item() not supported for this tensor type",
1071 "item",
1072 )),
1073 }
1074 }
1075
1076 pub fn get_scalar_i64(&self) -> Result<i64> {
1082 self.item::<i64>()
1083 }
1084
1085 pub fn eq_scalar(&self, scalar: f64) -> Result<Tensor> {
1096 match self {
1097 Tensor::F32(a) => {
1098 let scalar_f32 = scalar as f32;
1099 let result =
1100 a.mapv(|x| if (x - scalar_f32).abs() < 1e-6 { 1.0f32 } else { 0.0f32 });
1101 Ok(Tensor::F32(result))
1102 },
1103 Tensor::F64(a) => {
1104 let result = a.mapv(|x| if (x - scalar).abs() < 1e-9 { 1.0f64 } else { 0.0f64 });
1105 Ok(Tensor::F64(result))
1106 },
1107 Tensor::I64(a) => {
1108 let scalar_i64 = scalar as i64;
1109 let result = a.mapv(|x| if x == scalar_i64 { 1i64 } else { 0i64 });
1110 Ok(Tensor::I64(result))
1111 },
1112 _ => Err(TrustformersError::tensor_op_error(
1113 "eq_scalar not supported for this tensor type",
1114 "eq_scalar",
1115 )),
1116 }
1117 }
1118
1119 pub fn batch_split(&self, batch_size: usize) -> Result<Vec<Tensor>> {
1142 if batch_size == 0 {
1143 return Err(TrustformersError::tensor_op_error(
1144 "Batch size must be greater than 0",
1145 "batch_split",
1146 ));
1147 }
1148
1149 let shape = self.shape();
1150 if shape.is_empty() {
1151 return Err(TrustformersError::tensor_op_error(
1152 "Cannot batch split a scalar tensor",
1153 "batch_split",
1154 ));
1155 }
1156
1157 let total_size = shape[0];
1158 let mut batches = Vec::new();
1159
1160 for start in (0..total_size).step_by(batch_size) {
1161 let end = std::cmp::min(start + batch_size, total_size);
1162 let batch = self.slice(0, start, end)?;
1163 batches.push(batch);
1164 }
1165
1166 Ok(batches)
1167 }
1168
1169 pub fn batch_stack(tensors: &[&Tensor]) -> Result<Tensor> {
1192 if tensors.is_empty() {
1193 return Err(TrustformersError::tensor_op_error(
1194 "Cannot stack empty tensor list",
1195 "batch_stack",
1196 ));
1197 }
1198
1199 let reference_shape = tensors[0].shape();
1201 for (i, tensor) in tensors.iter().enumerate() {
1202 if tensor.shape() != reference_shape {
1203 return Err(TrustformersError::tensor_op_error(
1204 &format!(
1205 "Tensor {} has shape {:?}, expected {:?}",
1206 i,
1207 tensor.shape(),
1208 reference_shape
1209 ),
1210 "batch_stack",
1211 ));
1212 }
1213 }
1214
1215 let mut new_shape = vec![tensors.len()];
1217 new_shape.extend_from_slice(&reference_shape);
1218
1219 match tensors[0] {
1220 Tensor::F32(_) => {
1221 let mut result_data = Vec::new();
1222 for tensor in tensors {
1223 if let Tensor::F32(arr) = tensor {
1224 result_data.extend(arr.iter().copied());
1225 }
1226 }
1227 Tensor::from_vec(result_data, &new_shape)
1228 },
1229 _ => Err(TrustformersError::tensor_op_error(
1230 "Batch stacking currently only implemented for F32 tensors",
1231 "batch_stack",
1232 )),
1233 }
1234 }
1235
1236 pub fn unbatch(&self) -> Result<Vec<Tensor>> {
1253 let shape = self.shape();
1254 if shape.is_empty() {
1255 return Err(TrustformersError::tensor_op_error(
1256 "Cannot unbatch a scalar tensor",
1257 "unbatch",
1258 ));
1259 }
1260
1261 let batch_size = shape[0];
1262 let mut items = Vec::with_capacity(batch_size);
1263
1264 for i in 0..batch_size {
1265 let item = self.slice(0, i, i + 1)?;
1266 let squeezed = item.squeeze(0)?;
1268 items.push(squeezed);
1269 }
1270
1271 Ok(items)
1272 }
1273}
1274
1275#[cfg(test)]
1276mod tests {
1277 use super::*;
1278
1279 #[test]
1280 fn test_gradient_tracking_basic() {
1281 enable_grad();
1283
1284 let mut x = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1285 let grad = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3])
1286 .expect("Tensor from_vec failed");
1287
1288 assert!(x.set_grad(grad.clone()).is_ok());
1290
1291 let retrieved_grad = x.grad().expect("operation failed in test");
1293 assert_eq!(retrieved_grad.shape(), vec![2, 3]);
1294
1295 disable_grad();
1296 }
1297
1298 #[test]
1299 fn test_gradient_tracking_disabled() {
1300 disable_grad();
1302
1303 let mut x = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1304 let grad = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1305
1306 assert!(x.set_grad(grad).is_err());
1308 assert!(x.grad().is_err());
1309 }
1310
1311 #[test]
1312 fn test_gradient_shape_validation() {
1313 enable_grad();
1314
1315 let mut x = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1316 let wrong_shape_grad = Tensor::ones(&[3, 2]).expect("Failed to create ones tensor");
1317
1318 assert!(x.set_grad(wrong_shape_grad).is_err());
1320
1321 disable_grad();
1322 }
1323
1324 #[test]
1325 fn test_clear_gradients() {
1326 enable_grad();
1327
1328 let mut x = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1329 let grad = Tensor::ones(&[2, 3]).expect("Failed to create ones tensor");
1330
1331 x.set_grad(grad).expect("operation failed in test");
1333
1334 assert!(x.grad().is_ok());
1336
1337 clear_gradients();
1339
1340 assert!(x.grad().is_err());
1342
1343 disable_grad();
1344 }
1345
1346 #[test]
1347 fn test_gradient_mode_functions() {
1348 disable_grad();
1350 assert!(!is_grad_enabled());
1351
1352 enable_grad();
1353 assert!(is_grad_enabled());
1354
1355 disable_grad();
1356 assert!(!is_grad_enabled());
1357 }
1358
1359 #[test]
1360 fn test_softmax_entropy_normalized_bounds() {
1361 let uniform =
1363 Tensor::from_vec(vec![1.0, 1.0, 1.0, 1.0], &[4]).expect("failed to build tensor");
1364 let h_uniform = uniform.softmax_entropy_normalized().expect("entropy failed");
1365 assert!(
1366 h_uniform > 0.99,
1367 "uniform entropy should be ~1.0, got {h_uniform}"
1368 );
1369
1370 let peaked =
1372 Tensor::from_vec(vec![20.0, 0.0, 0.0, 0.0], &[4]).expect("failed to build tensor");
1373 let h_peaked = peaked.softmax_entropy_normalized().expect("entropy failed");
1374 assert!(
1375 (0.0..0.2).contains(&h_peaked),
1376 "peaked entropy should be small, got {h_peaked}"
1377 );
1378
1379 let single = Tensor::from_vec(vec![5.0], &[1]).expect("failed to build tensor");
1381 assert_eq!(
1382 single.softmax_entropy_normalized().expect("entropy failed"),
1383 0.0
1384 );
1385 }
1386}