use super::sum_mean::SegmentLayout;
use crate::tensor::TensorStorage;
use crate::{Result, Tensor, TensorError};
pub fn segment_prod<T>(
data: &Tensor<T>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<T>>
where
T: Clone
+ Default
+ std::ops::Mul<Output = T>
+ scirs2_core::num_traits::One
+ Send
+ Sync
+ 'static
+ bytemuck::Pod
+ bytemuck::Zeroable,
{
if data.shape().dims()[0] != segment_ids.shape().dims()[0] {
return Err(TensorError::shape_mismatch(
"segment_prod",
"data and segment_ids must have same first dimension",
&format!(
"data: {:?}, segment_ids: {:?}",
data.shape().dims(),
segment_ids.shape().dims()
),
));
}
let layout = SegmentLayout::new(data.shape().dims(), num_segments);
match (&data.storage, &segment_ids.storage) {
(TensorStorage::Cpu(data_arr), TensorStorage::Cpu(ids_arr)) => {
let data_flat = data_arr.iter().cloned().collect::<Vec<T>>();
let ids = ids_arr.iter().copied().collect::<Vec<i32>>();
let (result, _initialized) =
layout.reduce_rows(&data_flat, &ids, T::one(), |a, b| a.clone() * b.clone());
Tensor::from_vec(result, &layout.out_shape)
}
#[cfg(feature = "gpu")]
_ => {
let cpu_data = data.to_cpu()?;
let cpu_ids = segment_ids.to_cpu()?;
segment_prod(&cpu_data, &cpu_ids, num_segments)
}
}
}
pub fn segment_any(
data: &Tensor<u8>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<u8>> {
if data.shape().dims()[0] != segment_ids.shape().dims()[0] {
return Err(TensorError::shape_mismatch(
"segment_any",
"data and segment_ids must have same first dimension",
&format!(
"data: {:?}, segment_ids: {:?}",
data.shape().dims(),
segment_ids.shape().dims()
),
));
}
let layout = SegmentLayout::new(data.shape().dims(), num_segments);
match (&data.storage, &segment_ids.storage) {
(TensorStorage::Cpu(data_arr), TensorStorage::Cpu(ids_arr)) => {
let data_flat = data_arr.iter().copied().collect::<Vec<u8>>();
let ids = ids_arr.iter().copied().collect::<Vec<i32>>();
let (result, _initialized) =
layout.reduce_rows(&data_flat, &ids, 0u8, |a, b| u8::from(*a != 0 || *b != 0));
Tensor::from_vec(result, &layout.out_shape)
}
#[cfg(feature = "gpu")]
_ => {
let cpu_data = data.to_cpu()?;
let cpu_ids = segment_ids.to_cpu()?;
segment_any(&cpu_data, &cpu_ids, num_segments)
}
}
}
pub fn segment_all(
data: &Tensor<u8>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<u8>> {
if data.shape().dims()[0] != segment_ids.shape().dims()[0] {
return Err(TensorError::shape_mismatch(
"segment_all",
"data and segment_ids must have same first dimension",
&format!(
"data: {:?}, segment_ids: {:?}",
data.shape().dims(),
segment_ids.shape().dims()
),
));
}
let layout = SegmentLayout::new(data.shape().dims(), num_segments);
match (&data.storage, &segment_ids.storage) {
(TensorStorage::Cpu(data_arr), TensorStorage::Cpu(ids_arr)) => {
let data_flat = data_arr.iter().copied().collect::<Vec<u8>>();
let ids = ids_arr.iter().copied().collect::<Vec<i32>>();
let (result, _initialized) =
layout.reduce_rows(&data_flat, &ids, 1u8, |a, b| u8::from(*a != 0 && *b != 0));
Tensor::from_vec(result, &layout.out_shape)
}
#[cfg(feature = "gpu")]
_ => {
let cpu_data = data.to_cpu()?;
let cpu_ids = segment_ids.to_cpu()?;
segment_all(&cpu_data, &cpu_ids, num_segments)
}
}
}
#[cfg(all(test, feature = "gpu"))]
mod gpu_tests {
use super::*;
use crate::Device;
#[test]
fn gpu_segment_prod_matches_cpu_reference() {
let data_cpu = Tensor::<f32>::from_vec(vec![2.0, 3.0, 4.0, 5.0], &[4])
.expect("test: from_vec should succeed");
let ids_cpu =
Tensor::<i32>::from_vec(vec![0, 0, 1, 1], &[4]).expect("test: from_vec should succeed");
let (data_gpu, ids_gpu) = match (data_cpu.to(Device::Gpu(0)), ids_cpu.to(Device::Gpu(0))) {
(Ok(d), Ok(i)) => (d, i),
_ => return, };
let expected =
segment_prod(&data_cpu, &ids_cpu, 2).expect("test: CPU segment_prod should succeed");
let actual = segment_prod(&data_gpu, &ids_gpu, 2)
.expect("test: GPU segment_prod should succeed with a real adapter");
assert_eq!(actual.shape().dims(), expected.shape().dims());
assert_eq!(
actual.to_vec().expect("test: to_vec should succeed"),
expected.to_vec().expect("test: to_vec should succeed")
);
}
#[test]
fn gpu_segment_any_matches_cpu_reference() {
let data_cpu =
Tensor::<u8>::from_vec(vec![0, 1, 0, 0], &[4]).expect("test: from_vec should succeed");
let ids_cpu =
Tensor::<i32>::from_vec(vec![0, 0, 1, 1], &[4]).expect("test: from_vec should succeed");
let (data_gpu, ids_gpu) = match (data_cpu.to(Device::Gpu(0)), ids_cpu.to(Device::Gpu(0))) {
(Ok(d), Ok(i)) => (d, i),
_ => return, };
let expected =
segment_any(&data_cpu, &ids_cpu, 2).expect("test: CPU segment_any should succeed");
let actual = segment_any(&data_gpu, &ids_gpu, 2)
.expect("test: GPU segment_any should succeed with a real adapter");
assert_eq!(
actual.to_vec().expect("test: to_vec should succeed"),
expected.to_vec().expect("test: to_vec should succeed")
);
}
#[test]
fn gpu_segment_all_matches_cpu_reference() {
let data_cpu =
Tensor::<u8>::from_vec(vec![1, 1, 1, 0], &[4]).expect("test: from_vec should succeed");
let ids_cpu =
Tensor::<i32>::from_vec(vec![0, 0, 1, 1], &[4]).expect("test: from_vec should succeed");
let (data_gpu, ids_gpu) = match (data_cpu.to(Device::Gpu(0)), ids_cpu.to(Device::Gpu(0))) {
(Ok(d), Ok(i)) => (d, i),
_ => return, };
let expected =
segment_all(&data_cpu, &ids_cpu, 2).expect("test: CPU segment_all should succeed");
let actual = segment_all(&data_gpu, &ids_gpu, 2)
.expect("test: GPU segment_all should succeed with a real adapter");
assert_eq!(
actual.to_vec().expect("test: to_vec should succeed"),
expected.to_vec().expect("test: to_vec should succeed")
);
assert_eq!(
expected.to_vec().expect("test: to_vec should succeed"),
vec![1u8, 0u8]
);
}
}