1use candle_core::Tensor;
4use ferrum_interfaces::TensorLike;
5use ferrum_types::{DataType, Device, FerrumError, Result};
6use std::any::Any;
7
8#[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 pub fn from_tensorref(tensor_ref: &ferrum_interfaces::TensorRef) -> Option<Tensor> {
29 let _ = tensor_ref;
32
33 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 fn to_vec_f32(&self) -> Result<Vec<f32>> {
185 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 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 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 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 let squeezed = tensor
220 .squeeze(2)
221 .map_err(|e| FerrumError::model(format!("Squeeze dim 2 failed: {}", e)))?;
222
223 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 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 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 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 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}