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
76pub 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 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 tensor.dtype() == target_dtype {
129 return tensor;
130 }
131
132 let tensor = tensor.to_contiguous();
134 let shape = tensor.layout().shape().clone();
135
136 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 (tensor.dtype(), target_dtype) {
151 (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 (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 (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 (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 (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 (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 (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 (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 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 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 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 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}