use crate::tensor::TensorStorage;
use crate::{Result, Tensor, TensorError};
use rayon::prelude::*;
pub(super) struct SegmentLayout {
pub(super) num_segments: usize,
pub(super) feature_width: usize,
pub(super) out_shape: Vec<usize>,
}
impl SegmentLayout {
pub(super) fn new(data_dims: &[usize], num_segments: usize) -> Self {
let feature_width: usize = data_dims.iter().skip(1).product();
let mut out_shape = Vec::with_capacity(data_dims.len());
out_shape.push(num_segments);
out_shape.extend_from_slice(&data_dims[1..]);
Self {
num_segments,
feature_width,
out_shape,
}
}
fn row_count(&self, data_len: usize) -> usize {
data_len.checked_div(self.feature_width).unwrap_or(0)
}
fn add_row<T>(&self, acc: &mut [T], seg: usize, row: usize, data: &[T])
where
T: Clone + std::ops::Add<Output = T>,
{
let dst = seg * self.feature_width;
let src = row * self.feature_width;
for offset in 0..self.feature_width {
acc[dst + offset] = acc[dst + offset].clone() + data[src + offset].clone();
}
}
fn accumulate<T>(&self, data: &[T], ids: &[i32]) -> Vec<T>
where
T: Clone
+ Default
+ std::ops::Add<Output = T>
+ scirs2_core::num_traits::Zero
+ Send
+ Sync,
{
let rows = self.row_count(data.len());
let acc_len = self.num_segments * self.feature_width;
if rows > 1000 {
let chunk_size = std::cmp::max(1, rows / rayon::current_num_threads());
(0..rows)
.into_par_iter()
.chunks(chunk_size)
.map(|row_chunk| {
let mut local = vec![T::zero(); acc_len];
for row in row_chunk {
let seg = ids[row];
if seg >= 0 && (seg as usize) < self.num_segments {
self.add_row(&mut local, seg as usize, row, data);
}
}
local
})
.reduce(
|| vec![T::zero(); acc_len],
|mut a, b| {
for (slot, val) in a.iter_mut().zip(b) {
*slot = slot.clone() + val;
}
a
},
)
} else {
let mut acc = vec![T::zero(); acc_len];
for (row, &seg) in ids.iter().enumerate().take(rows) {
if seg >= 0 && (seg as usize) < self.num_segments {
self.add_row(&mut acc, seg as usize, row, data);
}
}
acc
}
}
fn accumulate_with_counts<T>(&self, data: &[T], ids: &[i32]) -> (Vec<T>, Vec<usize>)
where
T: Clone + Default + std::ops::Add<Output = T> + scirs2_core::num_traits::Zero,
{
let rows = self.row_count(data.len());
let acc_len = self.num_segments * self.feature_width;
let mut acc = vec![T::zero(); acc_len];
let mut counts = vec![0usize; self.num_segments];
for (row, &seg) in ids.iter().enumerate().take(rows) {
if seg >= 0 && (seg as usize) < self.num_segments {
let seg = seg as usize;
self.add_row(&mut acc, seg, row, data);
counts[seg] += 1;
}
}
(acc, counts)
}
fn combine_row<T, F>(
&self,
acc: &mut [T],
initialized: &mut [bool],
seg: usize,
row: usize,
data: &[T],
combine: &F,
) where
T: Clone,
F: Fn(&T, &T) -> T,
{
let dst = seg * self.feature_width;
let src = row * self.feature_width;
for offset in 0..self.feature_width {
let cell = dst + offset;
if initialized[cell] {
acc[cell] = combine(&acc[cell], &data[src + offset]);
} else {
acc[cell] = data[src + offset].clone();
initialized[cell] = true;
}
}
}
pub(super) fn reduce_rows<T, F>(
&self,
data: &[T],
ids: &[i32],
identity: T,
combine: F,
) -> (Vec<T>, Vec<bool>)
where
T: Clone + Send + Sync,
F: Fn(&T, &T) -> T + Sync + Send,
{
let rows = self.row_count(data.len());
let acc_len = self.num_segments * self.feature_width;
if rows > 1000 {
let chunk_size = std::cmp::max(1, rows / rayon::current_num_threads());
(0..rows)
.into_par_iter()
.chunks(chunk_size)
.map(|row_chunk| {
let mut local_values = vec![identity.clone(); acc_len];
let mut local_initialized = vec![false; acc_len];
for row in row_chunk {
let seg = ids[row];
if seg >= 0 && (seg as usize) < self.num_segments {
self.combine_row(
&mut local_values,
&mut local_initialized,
seg as usize,
row,
data,
&combine,
);
}
}
(local_values, local_initialized)
})
.reduce(
|| (vec![identity.clone(); acc_len], vec![false; acc_len]),
|mut a, b| {
merge_partial_rows(&mut a.0, &mut a.1, b, &combine);
a
},
)
} else {
let mut values = vec![identity.clone(); acc_len];
let mut initialized = vec![false; acc_len];
for (row, &seg) in ids.iter().enumerate().take(rows) {
if seg >= 0 && (seg as usize) < self.num_segments {
self.combine_row(
&mut values,
&mut initialized,
seg as usize,
row,
data,
&combine,
);
}
}
(values, initialized)
}
}
}
fn merge_partial_rows<T, F>(
dst_values: &mut [T],
dst_initialized: &mut [bool],
src: (Vec<T>, Vec<bool>),
combine: &F,
) where
T: Clone,
F: Fn(&T, &T) -> T,
{
let (src_values, src_initialized) = src;
for i in 0..dst_values.len() {
if src_initialized[i] {
if dst_initialized[i] {
dst_values[i] = combine(&dst_values[i], &src_values[i]);
} else {
dst_values[i] = src_values[i].clone();
dst_initialized[i] = true;
}
}
}
}
pub fn segment_sum<T>(
data: &Tensor<T>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<T>>
where
T: Clone
+ Default
+ std::ops::Add<Output = T>
+ scirs2_core::num_traits::Zero
+ 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 = layout.accumulate(&data_flat, &ids);
Tensor::from_vec(result, &layout.out_shape)
}
#[cfg(feature = "gpu")]
_ => {
let cpu_data = data.to_cpu()?;
let cpu_ids = segment_ids.to_cpu()?;
segment_sum(&cpu_data, &cpu_ids, num_segments)
}
}
}
pub fn segment_mean<T>(
data: &Tensor<T>,
segment_ids: &Tensor<i32>,
num_segments: usize,
) -> Result<Tensor<T>>
where
T: Clone
+ Default
+ std::ops::Add<Output = T>
+ std::ops::Div<Output = T>
+ scirs2_core::num_traits::Zero
+ scirs2_core::num_traits::FromPrimitive
+ 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 (mut result, counts) = layout.accumulate_with_counts(&data_flat, &ids);
for (segment, &count) in counts.iter().enumerate() {
if count > 0 {
if let Some(count_t) = T::from_usize(count) {
let base = segment * layout.feature_width;
for offset in 0..layout.feature_width {
let cell = base + offset;
result[cell] = result[cell].clone() / count_t.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_mean(&cpu_data, &cpu_ids, num_segments)
}
}
}
#[cfg(all(test, feature = "gpu"))]
mod gpu_tests {
use super::*;
use crate::Device;
#[test]
fn gpu_segment_sum_matches_cpu_reference() {
let data_cpu = Tensor::<f32>::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.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_sum(&data_cpu, &ids_cpu, 3).expect("test: CPU segment_sum should succeed");
let actual = segment_sum(&data_gpu, &ids_gpu, 3)
.expect("test: GPU segment_sum 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_mean_matches_cpu_reference() {
let data_cpu = Tensor::<f32>::from_vec(vec![2.0, 4.0, 10.0, 20.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_mean(&data_cpu, &ids_cpu, 2).expect("test: CPU segment_mean should succeed");
let actual = segment_mean(&data_gpu, &ids_gpu, 2)
.expect("test: GPU segment_mean should succeed with a real adapter");
assert_eq!(actual.shape().dims(), expected.shape().dims());
let actual_vals = actual.to_vec().expect("test: to_vec should succeed");
let expected_vals = expected.to_vec().expect("test: to_vec should succeed");
assert_eq!(expected_vals, vec![3.0, 15.0]);
for (a, e) in actual_vals.iter().zip(expected_vals.iter()) {
assert!(
(a - e).abs() < 1e-6,
"GPU segment_mean {a} does not match CPU reference {e} within tolerance"
);
}
}
}