use super::sum_mean::SegmentLayout;
use crate::tensor::TensorStorage;
use crate::{Result, Tensor, TensorError};
pub fn segment_max<T>(
data: &Tensor<T>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<T>>
where
T: Clone
+ Default
+ PartialOrd
+ scirs2_core::num_traits::Bounded
+ Send
+ Sync
+ 'static
+ bytemuck::Pod
+ bytemuck::Zeroable,
{
if data.shape().dims()[0] != segment_ids.shape().dims()[0] {
return Err(TensorError::shape_mismatch(
"segment_reduction",
"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::min_value(), |a, b| {
if b > a {
b.clone()
} else {
a.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_max(&cpu_data, &cpu_ids, num_segments)
}
}
}
pub fn segment_min<T>(
data: &Tensor<T>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<T>>
where
T: Clone
+ Default
+ PartialOrd
+ scirs2_core::num_traits::Bounded
+ Send
+ Sync
+ 'static
+ bytemuck::Pod
+ bytemuck::Zeroable,
{
if data.shape().dims()[0] != segment_ids.shape().dims()[0] {
return Err(TensorError::shape_mismatch(
"segment_reduction",
"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::max_value(), |a, b| {
if b < a {
b.clone()
} else {
a.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_min(&cpu_data, &cpu_ids, num_segments)
}
}
}
#[cfg(all(test, feature = "gpu"))]
mod gpu_tests {
use super::*;
use crate::Device;
#[test]
fn gpu_segment_max_matches_cpu_reference() {
let data_cpu = Tensor::<f32>::from_vec(vec![1.0, 5.0, 3.0, 2.0, 8.0, 0.0], &[6])
.expect("test: from_vec should succeed");
let ids_cpu = Tensor::<i32>::from_vec(vec![0, 0, 1, 1, 2, 2], &[6])
.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_max(&data_cpu, &ids_cpu, 3).expect("test: CPU segment_max should succeed");
let actual = segment_max(&data_gpu, &ids_gpu, 3)
.expect("test: GPU segment_max 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_min_matches_cpu_reference() {
let data_cpu = Tensor::<f32>::from_vec(vec![1.0, 5.0, 3.0, 2.0, 8.0, 0.0], &[6])
.expect("test: from_vec should succeed");
let ids_cpu = Tensor::<i32>::from_vec(vec![0, 0, 1, 1, 2, 2], &[6])
.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_min(&data_cpu, &ids_cpu, 3).expect("test: CPU segment_min should succeed");
let actual = segment_min(&data_gpu, &ids_gpu, 3)
.expect("test: GPU segment_min 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")
);
}
}