Skip to main content

ferrum_engine/tensor_factory/
candle.rs

1//! Candle backend - MVP implementation with core functionality
2//!
3//! Provides basic Candle tensor operations for CPU and GPU devices.
4
5use ferrum_interfaces::{TensorFactory, TensorLike, TensorOps, TensorRef};
6use ferrum_types::{DataType, Device, Result};
7use std::any::Any;
8use std::sync::Arc;
9
10/// Candle tensor wrapper
11pub 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        // Extract tensor data as Vec<f32>
81        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        // MVP: simplified, return clone
169        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        // Fast path for greedy sampling: compute argmax on the tensor's device,
182        // then transfer only a single scalar to CPU.
183        //
184        // This is intentionally conservative: it assumes batch=1 and returns the
185        // first element when batch exists.
186        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                // [batch, vocab] -> take batch 0 -> [vocab]
193                self.inner.i(0).map_err(|e| {
194                    ferrum_types::FerrumError::backend(format!("Index batch failed: {}", e))
195                })?
196            }
197            3 => {
198                // [batch, seq, vocab] -> take batch 0, last seq -> [vocab]
199                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        // Candle argmax returns a tensor; we read back a single u32.
213        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
229/// Candle tensor factory
230pub 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/// Candle tensor operations
366#[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        // MVP: simplified layer norm
441        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
534// ============================================================================
535// Helper Functions
536// ============================================================================
537
538fn get_candle_tensor(tensor: &TensorRef) -> Result<&candle_core::Tensor> {
539    // MVP: use type_id check since as_any not in TensorLike trait yet
540    let concrete_ref: &CandleTensor = unsafe {
541        // This is safe if we always create tensors through this backend
542        &*(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), // Fallback
555        _ => 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)), // Default to GPU 0
613        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// ============================================================================
629// Unit Tests
630// ============================================================================
631
632#[cfg(test)]
633mod tests {
634    use super::*;
635
636    #[test]
637    fn test_dtype_conversions() {
638        // FP32
639        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        // FP16
644        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        // 2x2 matrices
701        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}