1use ferrum_interfaces::{TensorFactory, TensorLike, TensorOps, TensorRef};
6use ferrum_types::{DataType, Device, Result};
7use std::any::Any;
8use std::sync::Arc;
9
10pub struct CandleTensor {
12 inner: candle_core::Tensor,
13 device: Device,
14 dtype: DataType,
15}
16
17impl CandleTensor {
18 pub fn new(tensor: candle_core::Tensor) -> Result<Self> {
19 let device = candle_device_to_ferrum(tensor.device())?;
20 let dtype = candle_dtype_to_ferrum(tensor.dtype())?;
21
22 Ok(Self {
23 inner: tensor,
24 device,
25 dtype,
26 })
27 }
28
29 pub fn inner(&self) -> &candle_core::Tensor {
30 &self.inner
31 }
32}
33
34impl std::fmt::Debug for CandleTensor {
35 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
36 f.debug_struct("CandleTensor")
37 .field("shape", &self.inner.dims())
38 .field("dtype", &self.dtype)
39 .field("device", &self.device)
40 .finish()
41 }
42}
43
44impl TensorLike for CandleTensor {
45 fn as_any(&self) -> &dyn Any {
46 self
47 }
48
49 fn shape(&self) -> &[usize] {
50 self.inner.dims()
51 }
52
53 fn dtype(&self) -> DataType {
54 self.dtype
55 }
56
57 fn device(&self) -> Device {
58 self.device.clone()
59 }
60
61 fn to_device(&self, device: &Device) -> Result<TensorRef> {
62 let candle_device = ferrum_device_to_candle(device.clone())?;
63 let moved = self
64 .inner
65 .to_device(&candle_device)
66 .map_err(|e| ferrum_types::FerrumError::backend(format!("Device transfer: {}", e)))?;
67 Ok(Arc::new(Self::new(moved)?))
68 }
69
70 fn to_dtype(&self, dtype: DataType) -> Result<TensorRef> {
71 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
72 let converted = self
73 .inner
74 .to_dtype(candle_dtype)
75 .map_err(|e| ferrum_types::FerrumError::backend(format!("DType conversion: {}", e)))?;
76 Ok(Arc::new(Self::new(converted)?))
77 }
78
79 fn to_vec_f32(&self) -> Result<Vec<f32>> {
80 match self.inner.dims().len() {
82 1 => self
83 .inner
84 .to_vec1::<f32>()
85 .map_err(|e| ferrum_types::FerrumError::backend(format!("to_vec1 failed: {}", e))),
86 2 => {
87 let batch = self.inner.to_vec2::<f32>().map_err(|e| {
88 ferrum_types::FerrumError::backend(format!("to_vec2 failed: {}", e))
89 })?;
90 Ok(batch.into_iter().next().unwrap_or_default())
91 }
92 3 => {
93 let all = self.inner.to_vec3::<f32>().map_err(|e| {
94 ferrum_types::FerrumError::backend(format!("to_vec3 failed: {}", e))
95 })?;
96 Ok(all
97 .into_iter()
98 .next()
99 .and_then(|seq| seq.into_iter().last())
100 .unwrap_or_default())
101 }
102 _ => Err(ferrum_types::FerrumError::backend(format!(
103 "Unsupported tensor dimensions: {:?}",
104 self.inner.dims()
105 ))),
106 }
107 }
108
109 fn to_vec_u32(&self) -> Result<Vec<u32>> {
110 let cpu_tensor = self
111 .inner
112 .to_device(&candle_core::Device::Cpu)
113 .map_err(|e| ferrum_types::FerrumError::backend(format!("to_cpu failed: {}", e)))?;
114
115 match cpu_tensor.dims().len() {
116 1 => match cpu_tensor.to_vec1::<u32>() {
117 Ok(tokens) => Ok(tokens),
118 Err(_) => cpu_tensor
119 .to_vec1::<f32>()
120 .map(|tokens| tokens.into_iter().map(|x| x as u32).collect())
121 .map_err(|e| {
122 ferrum_types::FerrumError::backend(format!(
123 "to_vec1<u32/f32> failed: {}",
124 e
125 ))
126 }),
127 },
128 2 => match cpu_tensor.to_vec2::<u32>() {
129 Ok(batch) => Ok(batch.into_iter().next().unwrap_or_default()),
130 Err(_) => cpu_tensor
131 .to_vec2::<f32>()
132 .map(|batch| {
133 batch
134 .into_iter()
135 .next()
136 .unwrap_or_default()
137 .into_iter()
138 .map(|x| x as u32)
139 .collect()
140 })
141 .map_err(|e| {
142 ferrum_types::FerrumError::backend(format!(
143 "to_vec2<u32/f32> failed: {}",
144 e
145 ))
146 }),
147 },
148 _ => Err(ferrum_types::FerrumError::backend(format!(
149 "Unsupported tensor dimensions for token extraction: {:?}",
150 cpu_tensor.dims()
151 ))),
152 }
153 }
154
155 fn reshape(&self, shape: &[usize]) -> Result<TensorRef> {
156 let reshaped = self
157 .inner
158 .reshape(shape)
159 .map_err(|e| ferrum_types::FerrumError::backend(format!("Reshape: {}", e)))?;
160 Ok(Arc::new(Self::new(reshaped)?))
161 }
162
163 fn to_cpu(&self) -> Result<TensorRef> {
164 self.to_device(&Device::CPU)
165 }
166
167 fn view(&self, _start: &[usize], _end: &[usize]) -> Result<TensorRef> {
168 Ok(Arc::new(Self {
170 inner: self.inner.clone(),
171 device: self.device.clone(),
172 dtype: self.dtype,
173 }))
174 }
175
176 fn is_contiguous(&self) -> bool {
177 self.inner.is_contiguous()
178 }
179
180 fn argmax_last_dim_u32(&self) -> Result<u32> {
181 use candle_core::{IndexOp, D};
187
188 let dims = self.inner.dims();
189 let logits_1d = match dims.len() {
190 1 => self.inner.clone(),
191 2 => {
192 self.inner.i(0).map_err(|e| {
194 ferrum_types::FerrumError::backend(format!("Index batch failed: {}", e))
195 })?
196 }
197 3 => {
198 let seq_len = dims[1];
200 self.inner.i((0, seq_len.saturating_sub(1))).map_err(|e| {
201 ferrum_types::FerrumError::backend(format!("Index last token failed: {}", e))
202 })?
203 }
204 _ => {
205 return Err(ferrum_types::FerrumError::backend(format!(
206 "argmax_last_dim_u32 unsupported dims: {:?}",
207 dims
208 )))
209 }
210 };
211
212 let idx = logits_1d
214 .argmax(D::Minus1)
215 .map_err(|e| ferrum_types::FerrumError::backend(format!("Argmax failed: {}", e)))?
216 .to_device(&candle_core::Device::Cpu)
217 .map_err(|e| {
218 ferrum_types::FerrumError::backend(format!("Argmax to CPU failed: {}", e))
219 })?
220 .to_vec0::<u32>()
221 .map_err(|e| {
222 ferrum_types::FerrumError::backend(format!("Argmax readback failed: {}", e))
223 })?;
224
225 Ok(idx)
226 }
227}
228
229pub struct CandleTensorFactory {
231 device: Device,
232}
233
234impl CandleTensorFactory {
235 pub fn new(device: Device) -> Self {
236 Self { device }
237 }
238}
239
240impl std::fmt::Debug for CandleTensorFactory {
241 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
242 f.debug_struct("CandleTensorFactory")
243 .field("device", &self.device)
244 .finish()
245 }
246}
247
248impl TensorFactory for CandleTensorFactory {
249 fn empty(&self, shape: &[usize], dtype: DataType, device: Device) -> Result<TensorRef> {
250 let candle_device = ferrum_device_to_candle(device)?;
251 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
252
253 let tensor = candle_core::Tensor::zeros(shape, candle_dtype, &candle_device)
254 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
255 Ok(Arc::new(CandleTensor::new(tensor)?))
256 }
257
258 fn from_slice(
259 &self,
260 data: &[f32],
261 shape: &[usize],
262 dtype: DataType,
263 device: Device,
264 ) -> Result<TensorRef> {
265 let candle_device = ferrum_device_to_candle(device)?;
266 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
267
268 let tensor = candle_core::Tensor::from_slice(data, shape, &candle_device)
269 .and_then(|t| t.to_dtype(candle_dtype))
270 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
271
272 Ok(Arc::new(CandleTensor::new(tensor)?))
273 }
274
275 fn to_device(&self, tensor: &TensorRef, device: Device) -> Result<TensorRef> {
276 tensor.to_device(&device)
277 }
278
279 fn narrow(
280 &self,
281 tensor: &TensorRef,
282 dim: usize,
283 start: usize,
284 length: usize,
285 ) -> Result<TensorRef> {
286 let candle_tensor = get_candle_tensor(tensor)?;
287 let narrowed = candle_tensor
288 .narrow(dim, start, length)
289 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
290 Ok(Arc::new(CandleTensor::new(narrowed)?))
291 }
292
293 fn reshape(&self, tensor: &TensorRef, shape: &[usize]) -> Result<TensorRef> {
294 tensor.reshape(shape)
295 }
296
297 fn zeros_like(&self, tensor: &TensorRef) -> Result<TensorRef> {
298 let candle_tensor = get_candle_tensor(tensor)?;
299 let zeros = candle_core::Tensor::zeros(
300 candle_tensor.shape(),
301 candle_tensor.dtype(),
302 candle_tensor.device(),
303 )
304 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
305 Ok(Arc::new(CandleTensor::new(zeros)?))
306 }
307
308 fn zeros(&self, shape: &[usize], dtype: DataType, device: &Device) -> Result<TensorRef> {
309 let candle_device = ferrum_device_to_candle(device.clone())?;
310 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
311
312 let tensor = candle_core::Tensor::zeros(shape, candle_dtype, &candle_device)
313 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
314 Ok(Arc::new(CandleTensor::new(tensor)?))
315 }
316
317 fn ones(&self, shape: &[usize], dtype: DataType, device: &Device) -> Result<TensorRef> {
318 let candle_device = ferrum_device_to_candle(device.clone())?;
319 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
320
321 let tensor = candle_core::Tensor::ones(shape, candle_dtype, &candle_device)
322 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
323 Ok(Arc::new(CandleTensor::new(tensor)?))
324 }
325
326 fn uniform(
327 &self,
328 shape: &[usize],
329 low: f32,
330 high: f32,
331 dtype: DataType,
332 device: &Device,
333 ) -> Result<TensorRef> {
334 let candle_device = ferrum_device_to_candle(device.clone())?;
335 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
336
337 let tensor = candle_core::Tensor::rand(low, high, shape, &candle_device)
338 .and_then(|t| t.to_dtype(candle_dtype))
339 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
340 Ok(Arc::new(CandleTensor::new(tensor)?))
341 }
342
343 fn normal(
344 &self,
345 shape: &[usize],
346 mean: f32,
347 std: f32,
348 dtype: DataType,
349 device: &Device,
350 ) -> Result<TensorRef> {
351 let candle_device = ferrum_device_to_candle(device.clone())?;
352 let candle_dtype = ferrum_dtype_to_candle(dtype)?;
353
354 let tensor = candle_core::Tensor::randn(mean, std, shape, &candle_device)
355 .and_then(|t| t.to_dtype(candle_dtype))
356 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
357 Ok(Arc::new(CandleTensor::new(tensor)?))
358 }
359
360 fn from_tensor(&self, tensor: &TensorRef, device: &Device) -> Result<TensorRef> {
361 tensor.to_device(device)
362 }
363}
364
365#[derive(Debug, Clone, Default)]
367pub struct CandleTensorOps;
368
369impl TensorOps for CandleTensorOps {
370 fn matmul(&self, a: &TensorRef, b: &TensorRef) -> Result<TensorRef> {
371 let a_candle = get_candle_tensor(a)?;
372 let b_candle = get_candle_tensor(b)?;
373
374 let result = a_candle
375 .matmul(b_candle)
376 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
377 Ok(Arc::new(CandleTensor::new(result)?))
378 }
379
380 fn add(&self, a: &TensorRef, b: &TensorRef) -> Result<TensorRef> {
381 let a_candle = get_candle_tensor(a)?;
382 let b_candle = get_candle_tensor(b)?;
383
384 let result =
385 (a_candle + b_candle).map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
386 Ok(Arc::new(CandleTensor::new(result)?))
387 }
388
389 fn mul(&self, a: &TensorRef, b: &TensorRef) -> Result<TensorRef> {
390 let a_candle = get_candle_tensor(a)?;
391 let b_candle = get_candle_tensor(b)?;
392
393 let result =
394 (a_candle * b_candle).map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
395 Ok(Arc::new(CandleTensor::new(result)?))
396 }
397
398 fn sub(&self, a: &TensorRef, b: &TensorRef) -> Result<TensorRef> {
399 let a_candle = get_candle_tensor(a)?;
400 let b_candle = get_candle_tensor(b)?;
401
402 let result =
403 (a_candle - b_candle).map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
404 Ok(Arc::new(CandleTensor::new(result)?))
405 }
406
407 fn div(&self, a: &TensorRef, b: &TensorRef) -> Result<TensorRef> {
408 let a_candle = get_candle_tensor(a)?;
409 let b_candle = get_candle_tensor(b)?;
410
411 let result =
412 (a_candle / b_candle).map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
413 Ok(Arc::new(CandleTensor::new(result)?))
414 }
415
416 fn softmax(&self, tensor: &TensorRef, dim: i32) -> Result<TensorRef> {
417 let candle_tensor = get_candle_tensor(tensor)?;
418 let dim_usize = if dim < 0 {
419 (candle_tensor.rank() as i32 + dim) as usize
420 } else {
421 dim as usize
422 };
423
424 let result = candle_nn::ops::softmax(candle_tensor, dim_usize)
425 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
426 Ok(Arc::new(CandleTensor::new(result)?))
427 }
428
429 fn layer_norm(
430 &self,
431 input: &TensorRef,
432 weight: &TensorRef,
433 bias: Option<&TensorRef>,
434 eps: f32,
435 ) -> Result<TensorRef> {
436 let input_candle = get_candle_tensor(input)?;
437 let weight_candle = get_candle_tensor(weight)?;
438 let _bias_candle = bias.map(|b| get_candle_tensor(b)).transpose()?;
439
440 let zero_bias = candle_core::Tensor::zeros(
442 weight_candle.shape(),
443 weight_candle.dtype(),
444 weight_candle.device(),
445 )
446 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
447
448 let bias_tensor = if let Some(b) = _bias_candle {
449 b
450 } else {
451 &zero_bias
452 };
453
454 let normalized = candle_nn::ops::layer_norm(input_candle, weight_candle, bias_tensor, eps)
455 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
456 Ok(Arc::new(CandleTensor::new(normalized)?))
457 }
458
459 fn rms_norm(&self, input: &TensorRef, weight: &TensorRef, eps: f32) -> Result<TensorRef> {
460 let input_candle = get_candle_tensor(input)?;
461 let weight_candle = get_candle_tensor(weight)?;
462
463 let _rms = candle_nn::RmsNorm::new(weight_candle.clone(), eps as f64);
464 let result = candle_nn::ops::rms_norm(input_candle, weight_candle, eps)
465 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
466 Ok(Arc::new(CandleTensor::new(result)?))
467 }
468
469 fn relu(&self, tensor: &TensorRef) -> Result<TensorRef> {
470 let candle_tensor = get_candle_tensor(tensor)?;
471 let result = candle_tensor
472 .relu()
473 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
474 Ok(Arc::new(CandleTensor::new(result)?))
475 }
476
477 fn gelu(&self, tensor: &TensorRef) -> Result<TensorRef> {
478 let candle_tensor = get_candle_tensor(tensor)?;
479 let result = candle_tensor
480 .gelu()
481 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
482 Ok(Arc::new(CandleTensor::new(result)?))
483 }
484
485 fn silu(&self, tensor: &TensorRef) -> Result<TensorRef> {
486 let candle_tensor = get_candle_tensor(tensor)?;
487 let result = candle_nn::ops::silu(candle_tensor)
488 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
489 Ok(Arc::new(CandleTensor::new(result)?))
490 }
491
492 fn concat(&self, tensors: &[&TensorRef], dim: usize) -> Result<TensorRef> {
493 let candle_tensors: Result<Vec<_>> = tensors.iter().map(|t| get_candle_tensor(t)).collect();
494 let candle_tensors = candle_tensors?;
495
496 let result = candle_core::Tensor::cat(&candle_tensors, dim)
497 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
498 Ok(Arc::new(CandleTensor::new(result)?))
499 }
500
501 fn split(&self, tensor: &TensorRef, sizes: &[usize], dim: usize) -> Result<Vec<TensorRef>> {
502 let candle_tensor = get_candle_tensor(tensor)?;
503 let mut result = Vec::new();
504 let mut offset = 0;
505
506 for &size in sizes {
507 let chunk = candle_tensor
508 .narrow(dim, offset, size)
509 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
510 result.push(Arc::new(CandleTensor::new(chunk)?) as TensorRef);
511 offset += size;
512 }
513
514 Ok(result)
515 }
516
517 fn transpose(&self, tensor: &TensorRef, dim0: usize, dim1: usize) -> Result<TensorRef> {
518 let candle_tensor = get_candle_tensor(tensor)?;
519 let result = candle_tensor
520 .transpose(dim0, dim1)
521 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
522 Ok(Arc::new(CandleTensor::new(result)?))
523 }
524
525 fn permute(&self, tensor: &TensorRef, dims: &[usize]) -> Result<TensorRef> {
526 let candle_tensor = get_candle_tensor(tensor)?;
527 let result = candle_tensor
528 .permute(dims)
529 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))?;
530 Ok(Arc::new(CandleTensor::new(result)?))
531 }
532}
533
534fn get_candle_tensor(tensor: &TensorRef) -> Result<&candle_core::Tensor> {
539 let concrete_ref: &CandleTensor = unsafe {
541 &*(Arc::as_ptr(tensor) as *const CandleTensor)
543 };
544 Ok(&concrete_ref.inner)
545}
546
547fn ferrum_dtype_to_candle(dtype: DataType) -> Result<candle_core::DType> {
548 match dtype {
549 DataType::FP32 => Ok(candle_core::DType::F32),
550 DataType::FP16 => Ok(candle_core::DType::F16),
551 DataType::BF16 => Ok(candle_core::DType::BF16),
552 DataType::UINT32 => Ok(candle_core::DType::U32),
553 DataType::UINT8 => Ok(candle_core::DType::U8),
554 DataType::INT32 => Ok(candle_core::DType::U32), _ => Err(ferrum_types::FerrumError::backend(format!(
556 "Unsupported dtype: {:?}",
557 dtype
558 ))),
559 }
560}
561
562fn candle_dtype_to_ferrum(dtype: candle_core::DType) -> Result<DataType> {
563 match dtype {
564 candle_core::DType::F32 => Ok(DataType::FP32),
565 candle_core::DType::F16 => Ok(DataType::FP16),
566 candle_core::DType::BF16 => Ok(DataType::BF16),
567 candle_core::DType::U32 => Ok(DataType::UINT32),
568 candle_core::DType::U8 => Ok(DataType::UINT8),
569 _ => Err(ferrum_types::FerrumError::backend(format!(
570 "Unsupported Candle dtype: {:?}",
571 dtype
572 ))),
573 }
574}
575
576fn ferrum_device_to_candle(device: Device) -> Result<candle_core::Device> {
577 match device {
578 Device::CPU => Ok(candle_core::Device::Cpu),
579 Device::CUDA(id) => {
580 #[cfg(feature = "candle-cuda-compat")]
581 {
582 candle_core::Device::new_cuda(id as usize)
583 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))
584 }
585 #[cfg(not(feature = "candle-cuda-compat"))]
586 {
587 let _ = id;
588 Err(ferrum_types::FerrumError::unsupported(
589 "legacy Candle CUDA tensors require the candle-cuda-compat feature",
590 ))
591 }
592 }
593 #[cfg(any(target_os = "macos", target_os = "ios"))]
594 Device::Metal => {
595 #[cfg(feature = "metal")]
596 {
597 candle_core::Device::new_metal(0)
598 .map_err(|e| ferrum_types::FerrumError::backend(e.to_string()))
599 }
600 #[cfg(not(feature = "metal"))]
601 {
602 Err(ferrum_types::FerrumError::unsupported("Metal not enabled"))
603 }
604 }
605 Device::ROCm(_) => Err(ferrum_types::FerrumError::unsupported("ROCm not supported")),
606 }
607}
608
609fn candle_device_to_ferrum(device: &candle_core::Device) -> Result<Device> {
610 match device {
611 candle_core::Device::Cpu => Ok(Device::CPU),
612 candle_core::Device::Cuda(_) => Ok(Device::CUDA(0)), candle_core::Device::Metal(_) => {
614 #[cfg(any(target_os = "macos", target_os = "ios"))]
615 {
616 Ok(Device::Metal)
617 }
618 #[cfg(not(any(target_os = "macos", target_os = "ios")))]
619 {
620 Err(ferrum_types::FerrumError::unsupported(
621 "Metal devices are not available on this platform",
622 ))
623 }
624 }
625 }
626}
627
628#[cfg(test)]
633mod tests {
634 use super::*;
635
636 #[test]
637 fn test_dtype_conversions() {
638 let candle_fp32 = ferrum_dtype_to_candle(DataType::FP32).unwrap();
640 let back_fp32 = candle_dtype_to_ferrum(candle_fp32).unwrap();
641 assert_eq!(back_fp32, DataType::FP32);
642
643 let candle_fp16 = ferrum_dtype_to_candle(DataType::FP16).unwrap();
645 let back_fp16 = candle_dtype_to_ferrum(candle_fp16).unwrap();
646 assert_eq!(back_fp16, DataType::FP16);
647 }
648
649 #[test]
650 fn test_device_conversions_cpu() {
651 let ferrum_device = Device::CPU;
652 let candle_device = ferrum_device_to_candle(ferrum_device.clone()).unwrap();
653 let back_device = candle_device_to_ferrum(&candle_device).unwrap();
654 assert_eq!(back_device, Device::CPU);
655 }
656
657 #[test]
658 fn test_tensor_factory_zeros() {
659 let factory = CandleTensorFactory::new(Device::CPU);
660 let tensor = factory
661 .zeros(&[2, 3], DataType::FP32, &Device::CPU)
662 .unwrap();
663
664 assert_eq!(tensor.shape(), &[2, 3]);
665 assert_eq!(tensor.dtype(), DataType::FP32);
666 }
667
668 #[test]
669 fn test_tensor_factory_ones() {
670 let factory = CandleTensorFactory::new(Device::CPU);
671 let tensor = factory.ones(&[2, 2], DataType::FP32, &Device::CPU).unwrap();
672
673 assert_eq!(tensor.shape(), &[2, 2]);
674 }
675
676 #[test]
677 fn test_tensor_ops_add() {
678 let factory = CandleTensorFactory::new(Device::CPU);
679 let ops = CandleTensorOps;
680
681 let a = factory
682 .from_slice(&[1.0, 2.0], &[2], DataType::FP32, Device::CPU)
683 .unwrap();
684 let b = factory
685 .from_slice(&[3.0, 4.0], &[2], DataType::FP32, Device::CPU)
686 .unwrap();
687
688 let result = ops.add(&a, &b).unwrap();
689 let data = result.to_vec_f32().unwrap();
690
691 assert!((data[0] - 4.0).abs() < 1e-5);
692 assert!((data[1] - 6.0).abs() < 1e-5);
693 }
694
695 #[test]
696 fn test_tensor_ops_matmul() {
697 let factory = CandleTensorFactory::new(Device::CPU);
698 let ops = CandleTensorOps;
699
700 let a = factory
702 .from_slice(&[1.0, 2.0, 3.0, 4.0], &[2, 2], DataType::FP32, Device::CPU)
703 .unwrap();
704 let b = factory
705 .from_slice(&[1.0, 0.0, 0.0, 1.0], &[2, 2], DataType::FP32, Device::CPU)
706 .unwrap();
707
708 let result = ops.matmul(&a, &b).unwrap();
709 assert_eq!(result.shape(), &[2, 2]);
710 }
711
712 #[test]
713 fn test_tensor_reshape() {
714 let factory = CandleTensorFactory::new(Device::CPU);
715 let tensor = factory
716 .zeros(&[2, 3], DataType::FP32, &Device::CPU)
717 .unwrap();
718
719 let reshaped = tensor.reshape(&[3, 2]).unwrap();
720 assert_eq!(reshaped.shape(), &[3, 2]);
721 }
722
723 #[test]
724 fn test_tensor_to_cpu() {
725 let factory = CandleTensorFactory::new(Device::CPU);
726 let tensor = factory
727 .zeros(&[2, 3], DataType::FP32, &Device::CPU)
728 .unwrap();
729
730 let cpu_tensor = tensor.to_cpu().unwrap();
731 assert_eq!(cpu_tensor.device(), Device::CPU);
732 }
733
734 #[test]
735 fn test_tensor_to_vec_u32_from_fp32_ids() {
736 let factory = CandleTensorFactory::new(Device::CPU);
737 let tensor = factory
738 .from_slice(&[1.0, 2.0, 3.0], &[1, 3], DataType::FP32, Device::CPU)
739 .unwrap();
740
741 let tokens = tensor.to_vec_u32().unwrap();
742 assert_eq!(tokens, vec![1, 2, 3]);
743 }
744}