Skip to main content

ruprim_host/
cast.rs

1use alloc::vec::Vec;
2use half::{bf16, f16};
3use ruda_core::{bytes::Bytes, tensor::{DType, FloatDType, IntDType, host::{HostTensor, Layout}}};
4
5pub fn bool_into_int(tensor: HostTensor, out_dtype: ruda_core::tensor::IntDType) -> HostTensor {
6    let tensor = tensor.to_contiguous();
7    let shape = tensor.layout().shape().clone();
8    let out_dt = DType::from(out_dtype);
9    let bools = tensor.bytes();
10
11    macro_rules! convert {
12        ($int_ty:ty) => {{
13            let data: Vec<$int_ty> =
14                bools.iter().map(|&x| if x != 0 { 1 } else { 0 }).collect();
15            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
16        }};
17    }
18
19    match out_dtype {
20        IntDType::I64 => convert!(i64),
21        IntDType::I32 => convert!(i32),
22        IntDType::I16 => convert!(i16),
23        IntDType::I8 => convert!(i8),
24        IntDType::U64 => convert!(u64),
25        IntDType::U32 => convert!(u32),
26        IntDType::U16 => convert!(u16),
27        IntDType::U8 => convert!(u8),
28    }
29}
30
31pub fn bool_into_float(
32    tensor: HostTensor,
33    out_dtype: ruda_core::tensor::FloatDType,
34) -> HostTensor {
35    let tensor = tensor.to_contiguous();
36    let shape = tensor.layout().shape().clone();
37    let out_dt = DType::from(out_dtype);
38    let bools = tensor.bytes();
39
40    match out_dtype {
41        FloatDType::F64 => {
42            let data: Vec<f64> = bools
43                .iter()
44                .map(|&x| if x != 0 { 1.0 } else { 0.0 })
45                .collect();
46            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
47        }
48        FloatDType::F32 | FloatDType::Flex32 => {
49            let data: Vec<f32> = bools
50                .iter()
51                .map(|&x| if x != 0 { 1.0 } else { 0.0 })
52                .collect();
53            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
54        }
55        FloatDType::F16 => {
56            let one = f16::from_f32(1.0);
57            let zero = f16::from_f32(0.0);
58            let data: Vec<f16> = bools
59                .iter()
60                .map(|&x| if x != 0 { one } else { zero })
61                .collect();
62            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
63        }
64        FloatDType::BF16 => {
65            let one = bf16::from_f32(1.0);
66            let zero = bf16::from_f32(0.0);
67            let data: Vec<bf16> = bools
68                .iter()
69                .map(|&x| if x != 0 { one } else { zero })
70                .collect();
71            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
72        }
73    }
74}
75
76// Precision limits: i64/u64 > 2^24 for f32/f16/bf16, > 2^53 for f64.
77pub fn int_into_float(
78    tensor: HostTensor,
79    out_dtype: ruda_core::tensor::FloatDType,
80) -> HostTensor {
81    let tensor = tensor.to_contiguous();
82    let shape = tensor.layout().shape().clone();
83    let src = tensor.dtype();
84    let out_dt = DType::from(out_dtype);
85
86    // Read source ints, applying conversion per-element.
87    // Each arm binds `$x` to the native int value; `$conv` must work for all int types.
88    macro_rules! read_ints {
89        (|$x:ident| $conv:expr) => {
90            match src {
91                DType::I64 => tensor.storage::<i64>().iter().map(|&$x| $conv).collect(),
92                DType::I32 => tensor.storage::<i32>().iter().map(|&$x| $conv).collect(),
93                DType::I16 => tensor.storage::<i16>().iter().map(|&$x| $conv).collect(),
94                DType::I8 => tensor.storage::<i8>().iter().map(|&$x| $conv).collect(),
95                DType::U64 => tensor.storage::<u64>().iter().map(|&$x| $conv).collect(),
96                DType::U32 => tensor.storage::<u32>().iter().map(|&$x| $conv).collect(),
97                DType::U16 => tensor.storage::<u16>().iter().map(|&$x| $conv).collect(),
98                DType::U8 => tensor.storage::<u8>().iter().map(|&$x| $conv).collect(),
99                _ => panic!("int_into_float: unsupported source dtype {:?}", src),
100            }
101        };
102    }
103
104    match out_dtype {
105        FloatDType::F64 => {
106            let data: Vec<f64> = read_ints!(|x| x as f64);
107            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
108        }
109        FloatDType::F32 | FloatDType::Flex32 => {
110            let data: Vec<f32> = read_ints!(|x| x as f32);
111            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
112        }
113        FloatDType::F16 => {
114            let data: Vec<f16> = read_ints!(|x| f16::from_f32(x as f32));
115            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
116        }
117        FloatDType::BF16 => {
118            let data: Vec<bf16> = read_ints!(|x| bf16::from_f32(x as f32));
119            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
120        }
121    }
122}
123
124pub fn int_cast(tensor: HostTensor, dtype: IntDType) -> HostTensor {
125    let target_dtype: DType = dtype.into();
126
127    // If already the target dtype, return as-is
128    if tensor.dtype() == target_dtype {
129        return tensor;
130    }
131
132    // Make contiguous for easier iteration
133    let tensor = tensor.to_contiguous();
134    let shape = tensor.layout().shape().clone();
135
136    // Helper macro to convert between types
137    macro_rules! cast_impl {
138        ($src_type:ty, $dst_type:ty, $dst_dtype:expr) => {{
139            let src: &[$src_type] = tensor.storage();
140            let dst: Vec<$dst_type> = src.iter().map(|&x| x as $dst_type).collect();
141            HostTensor::new(
142                Bytes::from_elems(dst),
143                Layout::contiguous(shape),
144                $dst_dtype,
145            )
146        }};
147    }
148
149    // Match source dtype to target dtype
150    match (tensor.dtype(), target_dtype) {
151        // From I64
152        (DType::I64, DType::I32) => cast_impl!(i64, i32, DType::I32),
153        (DType::I64, DType::I16) => cast_impl!(i64, i16, DType::I16),
154        (DType::I64, DType::I8) => cast_impl!(i64, i8, DType::I8),
155        (DType::I64, DType::U64) => cast_impl!(i64, u64, DType::U64),
156        (DType::I64, DType::U32) => cast_impl!(i64, u32, DType::U32),
157        (DType::I64, DType::U16) => cast_impl!(i64, u16, DType::U16),
158        (DType::I64, DType::U8) => cast_impl!(i64, u8, DType::U8),
159
160        // From I32
161        (DType::I32, DType::I64) => cast_impl!(i32, i64, DType::I64),
162        (DType::I32, DType::I16) => cast_impl!(i32, i16, DType::I16),
163        (DType::I32, DType::I8) => cast_impl!(i32, i8, DType::I8),
164        (DType::I32, DType::U64) => cast_impl!(i32, u64, DType::U64),
165        (DType::I32, DType::U32) => cast_impl!(i32, u32, DType::U32),
166        (DType::I32, DType::U16) => cast_impl!(i32, u16, DType::U16),
167        (DType::I32, DType::U8) => cast_impl!(i32, u8, DType::U8),
168
169        // From I16
170        (DType::I16, DType::I64) => cast_impl!(i16, i64, DType::I64),
171        (DType::I16, DType::I32) => cast_impl!(i16, i32, DType::I32),
172        (DType::I16, DType::I8) => cast_impl!(i16, i8, DType::I8),
173        (DType::I16, DType::U64) => cast_impl!(i16, u64, DType::U64),
174        (DType::I16, DType::U32) => cast_impl!(i16, u32, DType::U32),
175        (DType::I16, DType::U16) => cast_impl!(i16, u16, DType::U16),
176        (DType::I16, DType::U8) => cast_impl!(i16, u8, DType::U8),
177
178        // From I8
179        (DType::I8, DType::I64) => cast_impl!(i8, i64, DType::I64),
180        (DType::I8, DType::I32) => cast_impl!(i8, i32, DType::I32),
181        (DType::I8, DType::I16) => cast_impl!(i8, i16, DType::I16),
182        (DType::I8, DType::U64) => cast_impl!(i8, u64, DType::U64),
183        (DType::I8, DType::U32) => cast_impl!(i8, u32, DType::U32),
184        (DType::I8, DType::U16) => cast_impl!(i8, u16, DType::U16),
185        (DType::I8, DType::U8) => cast_impl!(i8, u8, DType::U8),
186
187        // From U64
188        (DType::U64, DType::I64) => cast_impl!(u64, i64, DType::I64),
189        (DType::U64, DType::I32) => cast_impl!(u64, i32, DType::I32),
190        (DType::U64, DType::I16) => cast_impl!(u64, i16, DType::I16),
191        (DType::U64, DType::I8) => cast_impl!(u64, i8, DType::I8),
192        (DType::U64, DType::U32) => cast_impl!(u64, u32, DType::U32),
193        (DType::U64, DType::U16) => cast_impl!(u64, u16, DType::U16),
194        (DType::U64, DType::U8) => cast_impl!(u64, u8, DType::U8),
195
196        // From U32
197        (DType::U32, DType::I64) => cast_impl!(u32, i64, DType::I64),
198        (DType::U32, DType::I32) => cast_impl!(u32, i32, DType::I32),
199        (DType::U32, DType::I16) => cast_impl!(u32, i16, DType::I16),
200        (DType::U32, DType::I8) => cast_impl!(u32, i8, DType::I8),
201        (DType::U32, DType::U64) => cast_impl!(u32, u64, DType::U64),
202        (DType::U32, DType::U16) => cast_impl!(u32, u16, DType::U16),
203        (DType::U32, DType::U8) => cast_impl!(u32, u8, DType::U8),
204
205        // From U16
206        (DType::U16, DType::I64) => cast_impl!(u16, i64, DType::I64),
207        (DType::U16, DType::I32) => cast_impl!(u16, i32, DType::I32),
208        (DType::U16, DType::I16) => cast_impl!(u16, i16, DType::I16),
209        (DType::U16, DType::I8) => cast_impl!(u16, i8, DType::I8),
210        (DType::U16, DType::U64) => cast_impl!(u16, u64, DType::U64),
211        (DType::U16, DType::U32) => cast_impl!(u16, u32, DType::U32),
212        (DType::U16, DType::U8) => cast_impl!(u16, u8, DType::U8),
213
214        // From U8
215        (DType::U8, DType::I64) => cast_impl!(u8, i64, DType::I64),
216        (DType::U8, DType::I32) => cast_impl!(u8, i32, DType::I32),
217        (DType::U8, DType::I16) => cast_impl!(u8, i16, DType::I16),
218        (DType::U8, DType::I8) => cast_impl!(u8, i8, DType::I8),
219        (DType::U8, DType::U64) => cast_impl!(u8, u64, DType::U64),
220        (DType::U8, DType::U32) => cast_impl!(u8, u32, DType::U32),
221        (DType::U8, DType::U16) => cast_impl!(u8, u16, DType::U16),
222
223        _ => panic!(
224            "int_cast: unsupported conversion from {:?} to {:?}",
225            tensor.dtype(),
226            target_dtype
227        ),
228    }
229}
230
231pub fn float_into_int(tensor: HostTensor, out_dtype: ruda_core::tensor::IntDType) -> HostTensor {
232    let tensor = tensor.to_contiguous();
233    let shape = tensor.layout().shape().clone();
234    let src = tensor.dtype();
235    let out_dt = DType::from(out_dtype);
236
237    // Read source floats as f64 (lossless for f32/f16/bf16).
238    macro_rules! read_floats {
239        (|$x:ident| $conv:expr) => {
240            match src {
241                DType::F32 => tensor
242                    .storage::<f32>()
243                    .iter()
244                    .map(|v| {
245                        let $x = *v as f64;
246                        $conv
247                    })
248                    .collect(),
249                DType::F64 => tensor
250                    .storage::<f64>()
251                    .iter()
252                    .map(|v| {
253                        let $x = *v;
254                        $conv
255                    })
256                    .collect(),
257                DType::F16 => tensor
258                    .storage::<f16>()
259                    .iter()
260                    .map(|v| {
261                        let $x = f32::from(*v) as f64;
262                        $conv
263                    })
264                    .collect(),
265                DType::BF16 => tensor
266                    .storage::<bf16>()
267                    .iter()
268                    .map(|v| {
269                        let $x = f32::from(*v) as f64;
270                        $conv
271                    })
272                    .collect(),
273                _ => panic!("float_into_int: unsupported source dtype {:?}", src),
274            }
275        };
276    }
277
278    macro_rules! convert {
279        ($int_ty:ty) => {{
280            let data: Vec<$int_ty> = read_floats!(|x| x as $int_ty);
281            HostTensor::new(Bytes::from_elems(data), Layout::contiguous(shape), out_dt)
282        }};
283    }
284
285    match out_dtype {
286        IntDType::I64 => convert!(i64),
287        IntDType::I32 => convert!(i32),
288        IntDType::I16 => convert!(i16),
289        IntDType::I8 => convert!(i8),
290        IntDType::U64 => convert!(u64),
291        IntDType::U32 => convert!(u32),
292        IntDType::U16 => convert!(u16),
293        IntDType::U8 => convert!(u8),
294    }
295}
296
297pub fn float_cast(tensor: HostTensor, dtype: FloatDType) -> HostTensor {
298    use ruda_core::tensor::host::Layout;
299    use ruda_core::bytes::Bytes;
300    use half::{bf16, f16};
301
302    let src_dtype = tensor.dtype();
303    let target_dtype = DType::from(dtype);
304
305    // No-op if already the same dtype
306    if src_dtype == target_dtype {
307        return tensor;
308    }
309
310    let tensor = tensor.to_contiguous();
311    let shape = tensor.layout().shape().clone();
312
313    // Convert to f64 intermediate, then to target
314    let f64_values: Vec<f64> = match src_dtype {
315        DType::F32 => {
316            let src: &[f32] = tensor.storage();
317            src.iter().map(|&v| v as f64).collect()
318        }
319        DType::F64 => {
320            let src: &[f64] = tensor.storage();
321            src.to_vec()
322        }
323        DType::F16 => {
324            let src: &[f16] = tensor.storage();
325            src.iter().map(|&v| v.to_f32() as f64).collect()
326        }
327        DType::BF16 => {
328            let src: &[bf16] = tensor.storage();
329            src.iter().map(|&v| v.to_f32() as f64).collect()
330        }
331        _ => panic!("float_cast: unsupported source dtype {:?}", src_dtype),
332    };
333
334    // Convert from f64 to target dtype
335    match target_dtype {
336        DType::F32 => {
337            let result: Vec<f32> = f64_values.iter().map(|&v| v as f32).collect();
338            let bytes = Bytes::from_elems(result);
339            HostTensor::new(bytes, Layout::contiguous(shape), DType::F32)
340        }
341        DType::F64 => {
342            let bytes = Bytes::from_elems(f64_values);
343            HostTensor::new(bytes, Layout::contiguous(shape), DType::F64)
344        }
345        DType::F16 => {
346            let result: Vec<f16> = f64_values.iter().map(|&v| f16::from_f64(v)).collect();
347            let bytes = Bytes::from_elems(result);
348            HostTensor::new(bytes, Layout::contiguous(shape), DType::F16)
349        }
350        DType::BF16 => {
351            let result: Vec<bf16> = f64_values.iter().map(|&v| bf16::from_f64(v)).collect();
352            let bytes = Bytes::from_elems(result);
353            HostTensor::new(bytes, Layout::contiguous(shape), DType::BF16)
354        }
355        _ => panic!("float_cast: unsupported target dtype {:?}", target_dtype),
356    }
357}