Skip to main content

ferrum_models/
tensor_wrapper.rs

1//! Candle Tensor wrapper implementing TensorLike
2
3use candle_core::Tensor;
4use ferrum_interfaces::TensorLike;
5use ferrum_types::{DataType, Device, FerrumError, Result};
6use std::any::Any;
7
8/// Wrapper for Candle Tensor to implement TensorLike
9#[derive(Debug, Clone)]
10pub struct CandleTensorWrapper {
11    tensor: Tensor,
12}
13
14impl CandleTensorWrapper {
15    pub fn new(tensor: Tensor) -> Self {
16        Self { tensor }
17    }
18
19    pub fn inner(&self) -> &Tensor {
20        &self.tensor
21    }
22
23    pub fn into_inner(self) -> Tensor {
24        self.tensor
25    }
26
27    /// Safe extraction from Arc<dyn TensorLike>
28    pub fn from_tensorref(tensor_ref: &ferrum_interfaces::TensorRef) -> Option<Tensor> {
29        // Try to extract by getting raw data and reconstructing
30        // This is safe because we only read immutable data
31        let _ = tensor_ref;
32
33        // For now, return None if not our wrapper
34        // A better approach would be to add a method to TensorLike to extract data
35        None
36    }
37}
38
39impl TensorLike for CandleTensorWrapper {
40    fn as_any(&self) -> &dyn Any {
41        self
42    }
43
44    fn shape(&self) -> &[usize] {
45        self.tensor.dims()
46    }
47
48    fn dtype(&self) -> DataType {
49        match self.tensor.dtype() {
50            candle_core::DType::F32 => DataType::FP32,
51            candle_core::DType::F16 => DataType::FP16,
52            candle_core::DType::BF16 => DataType::BF16,
53            _ => DataType::FP32,
54        }
55    }
56
57    fn device(&self) -> Device {
58        match self.tensor.device() {
59            candle_core::Device::Cpu => Device::CPU,
60            candle_core::Device::Cuda(_) => Device::CUDA(0),
61            candle_core::Device::Metal(_) => {
62                #[cfg(any(target_os = "macos", target_os = "ios"))]
63                return Device::Metal;
64                #[cfg(not(any(target_os = "macos", target_os = "ios")))]
65                Device::CPU
66            }
67        }
68    }
69
70    fn is_contiguous(&self) -> bool {
71        self.tensor.is_contiguous()
72    }
73
74    fn view(&self, start: &[usize], end: &[usize]) -> Result<ferrum_interfaces::TensorRef> {
75        if start.len() != end.len() || start.len() != self.tensor.dims().len() {
76            return Err(FerrumError::model(format!(
77                "Invalid view dimensions: start={:?}, end={:?}, shape={:?}",
78                start,
79                end,
80                self.tensor.dims()
81            )));
82        }
83
84        let mut view = self.tensor.clone();
85        for (dim, (&start_idx, &end_idx)) in start.iter().zip(end.iter()).enumerate() {
86            if end_idx < start_idx {
87                return Err(FerrumError::model(format!(
88                    "Invalid view range on dim {}: {}..{}",
89                    dim, start_idx, end_idx
90                )));
91            }
92
93            let current_dim = view
94                .dims()
95                .get(dim)
96                .copied()
97                .ok_or_else(|| FerrumError::model("View dimension out of bounds"))?;
98            if end_idx > current_dim {
99                return Err(FerrumError::model(format!(
100                    "View end out of bounds on dim {}: {} > {}",
101                    dim, end_idx, current_dim
102                )));
103            }
104
105            let length = end_idx - start_idx;
106            if start_idx != 0 || length != current_dim {
107                view = view.narrow(dim, start_idx, length).map_err(|e| {
108                    FerrumError::model(format!("View narrow failed on dim {}: {}", dim, e))
109                })?;
110            }
111        }
112
113        Ok(std::sync::Arc::new(CandleTensorWrapper::new(view)))
114    }
115
116    fn reshape(&self, shape: &[usize]) -> Result<ferrum_interfaces::TensorRef> {
117        let reshaped = self
118            .tensor
119            .reshape(shape)
120            .map_err(|e| FerrumError::model(format!("Reshape failed: {}", e)))?;
121        Ok(std::sync::Arc::new(CandleTensorWrapper::new(reshaped)))
122    }
123
124    fn to_cpu(&self) -> Result<ferrum_interfaces::TensorRef> {
125        if matches!(self.tensor.device(), candle_core::Device::Cpu) {
126            return Ok(std::sync::Arc::new(self.clone()));
127        }
128
129        let cpu_tensor = self
130            .tensor
131            .to_device(&candle_core::Device::Cpu)
132            .map_err(|e| FerrumError::model(format!("to_cpu failed: {}", e)))?;
133        Ok(std::sync::Arc::new(CandleTensorWrapper::new(cpu_tensor)))
134    }
135
136    fn to_device(&self, device: &Device) -> Result<ferrum_interfaces::TensorRef> {
137        let candle_device = match device {
138            Device::CPU => candle_core::Device::Cpu,
139            #[cfg(feature = "candle-cuda-compat")]
140            Device::CUDA(id) => candle_core::Device::new_cuda(*id)
141                .map_err(|e| FerrumError::device(format!("CUDA device error: {}", e)))?,
142            #[cfg(not(feature = "candle-cuda-compat"))]
143            Device::CUDA(_) => {
144                return Err(FerrumError::unsupported(
145                    "legacy Candle CUDA tensor transfer requires the candle-cuda-compat feature",
146                ));
147            }
148            #[cfg(any(target_os = "macos", target_os = "ios"))]
149            Device::Metal => candle_core::Device::new_metal(0)
150                .map_err(|e| FerrumError::device(format!("Metal device error: {}", e)))?,
151            Device::ROCm(_) => {
152                return Err(FerrumError::device("ROCm not supported yet"));
153            }
154        };
155
156        let device_tensor = self
157            .tensor
158            .to_device(&candle_device)
159            .map_err(|e| FerrumError::model(format!("to_device failed: {}", e)))?;
160        Ok(std::sync::Arc::new(CandleTensorWrapper::new(device_tensor)))
161    }
162
163    fn to_dtype(&self, dtype: DataType) -> Result<ferrum_interfaces::TensorRef> {
164        let candle_dtype = match &dtype {
165            DataType::FP32 => candle_core::DType::F32,
166            DataType::FP16 => candle_core::DType::F16,
167            DataType::BF16 => candle_core::DType::BF16,
168            _ => {
169                return Err(FerrumError::model(format!(
170                    "Unsupported dtype: {:?}",
171                    dtype
172                )))
173            }
174        };
175
176        let converted = self
177            .tensor
178            .to_dtype(candle_dtype)
179            .map_err(|e| FerrumError::model(format!("to_dtype failed: {}", e)))?;
180        Ok(std::sync::Arc::new(CandleTensorWrapper::new(converted)))
181    }
182
183    /// Extract tensor data as Vec<f32> - Candle implementation
184    fn to_vec_f32(&self) -> Result<Vec<f32>> {
185        // Ensure F32 dtype (CUDA/Metal may produce F16/BF16 logits)
186        let tensor = if self.tensor.dtype() != candle_core::DType::F32 {
187            self.tensor
188                .to_dtype(candle_core::DType::F32)
189                .map_err(|e| FerrumError::model(format!("Cast to f32 failed: {}", e)))?
190        } else {
191            self.tensor.clone()
192        };
193        // Handle different tensor dimensions
194        match tensor.dims().len() {
195            1 => tensor
196                .to_vec1::<f32>()
197                .map_err(|e| FerrumError::model(format!("to_vec1 failed: {}", e))),
198            2 => {
199                // Take first batch: [batch, vocab] -> [vocab]
200                let batch = tensor
201                    .to_vec2::<f32>()
202                    .map_err(|e| FerrumError::model(format!("to_vec2 failed: {}", e)))?;
203                Ok(batch.into_iter().next().unwrap_or_default())
204            }
205            3 => {
206                // Take last token of first batch: [batch, seq, vocab] -> [vocab]
207                let all = tensor
208                    .to_vec3::<f32>()
209                    .map_err(|e| FerrumError::model(format!("to_vec3 failed: {}", e)))?;
210                Ok(all
211                    .into_iter()
212                    .next()
213                    .and_then(|seq| seq.into_iter().last())
214                    .unwrap_or_default())
215            }
216            4 => {
217                // Handle [batch, seq, extra, vocab] - squeeze and take last
218                // First squeeze to 3D by selecting first element of extra dim
219                let squeezed = tensor
220                    .squeeze(2)
221                    .map_err(|e| FerrumError::model(format!("Squeeze dim 2 failed: {}", e)))?;
222
223                // Now extract as 3D: [batch, seq, vocab]
224                let all = squeezed
225                    .to_vec3::<f32>()
226                    .map_err(|e| FerrumError::model(format!("to_vec3 (from 4D) failed: {}", e)))?;
227                Ok(all
228                    .into_iter()
229                    .next()
230                    .and_then(|seq| seq.into_iter().last())
231                    .unwrap_or_default())
232            }
233            _ => Err(FerrumError::model(format!(
234                "Unsupported dims: {:?}",
235                self.tensor.dims()
236            ))),
237        }
238    }
239
240    fn to_vec_u32(&self) -> Result<Vec<u32>> {
241        // Handle different tensor dimensions for token IDs
242        match self.tensor.dims().len() {
243            1 => self
244                .tensor
245                .to_vec1::<u32>()
246                .map_err(|e| FerrumError::model(format!("to_vec1<u32> failed: {}", e))),
247            2 => {
248                // Take first batch: [batch, seq] -> [seq]
249                let batch = self
250                    .tensor
251                    .to_vec2::<u32>()
252                    .map_err(|e| FerrumError::model(format!("to_vec2<u32> failed: {}", e)))?;
253                Ok(batch.into_iter().next().unwrap_or_default())
254            }
255            _ => Err(FerrumError::model(format!(
256                "Unsupported dims for token extraction: {:?}",
257                self.tensor.dims()
258            ))),
259        }
260    }
261
262    fn argmax_last_dim_u32(&self) -> Result<u32> {
263        // Same strategy as runtime CandleTensor: argmax on-device, read back a scalar.
264        use candle_core::{IndexOp, D};
265
266        let dims = self.tensor.dims();
267        let logits_1d = match dims.len() {
268            1 => self.tensor.clone(),
269            2 => self
270                .tensor
271                .i(0)
272                .map_err(|e| FerrumError::model(format!("Index batch failed: {}", e)))?,
273            3 => {
274                let seq_len = dims[1];
275                self.tensor
276                    .i((0, seq_len.saturating_sub(1)))
277                    .map_err(|e| FerrumError::model(format!("Index last token failed: {}", e)))?
278            }
279            4 => {
280                // [batch, seq, extra, vocab] -> take batch 0, last seq, extra 0 -> [vocab]
281                let seq_len = dims[1];
282                self.tensor
283                    .i((0, seq_len.saturating_sub(1), 0))
284                    .map_err(|e| {
285                        FerrumError::model(format!("Index last token (4D) failed: {}", e))
286                    })?
287            }
288            _ => {
289                return Err(FerrumError::model(format!(
290                    "argmax_last_dim_u32 unsupported dims: {:?}",
291                    dims
292                )))
293            }
294        };
295
296        let idx = logits_1d
297            .argmax(D::Minus1)
298            .map_err(|e| FerrumError::model(format!("Argmax failed: {}", e)))?
299            .to_device(&candle_core::Device::Cpu)
300            .map_err(|e| FerrumError::model(format!("Argmax to CPU failed: {}", e)))?
301            .to_vec0::<u32>()
302            .map_err(|e| FerrumError::model(format!("Argmax readback failed: {}", e)))?;
303
304        Ok(idx)
305    }
306}
307
308#[cfg(test)]
309mod tests {
310    use super::*;
311
312    #[test]
313    fn view_extracts_last_sequence_slice() {
314        let tensor = Tensor::from_vec(
315            vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0],
316            (1, 2, 3),
317            &candle_core::Device::Cpu,
318        )
319        .expect("create tensor");
320        let wrapper = CandleTensorWrapper::new(tensor);
321
322        let view = wrapper.view(&[0, 1, 0], &[1, 2, 3]).expect("slice view");
323        assert_eq!(view.shape(), &[1, 1, 3]);
324        assert_eq!(view.to_vec_f32().expect("to_vec_f32"), vec![4.0, 5.0, 6.0]);
325    }
326}