use super::types::IndexedSample;
use candle_core::{Device, Tensor};
use rayon::prelude::*;
pub(crate) fn pack_at_indices<F>(
samples: &[IndexedSample],
sample_indices: &[usize],
k: usize,
target_device: &Device,
fill: F,
) -> anyhow::Result<Tensor>
where
F: Fn(usize, usize, usize, u32) -> f32 + Sync,
{
let n = sample_indices.len();
let mut buf = vec![0.0f32; n * k];
buf.par_chunks_mut(k)
.zip(sample_indices.par_iter())
.enumerate()
.for_each(|(row, (chunk, &si))| {
let s = &samples[si];
let take = s.indices.len().min(k);
for (kk, &feat) in s.indices[..take].iter().enumerate() {
chunk[kk] = fill(row, kk, si, feat);
}
});
Ok(Tensor::from_vec(buf, (n, k), target_device)?)
}
pub fn pack_indices_values(
samples: &[IndexedSample],
sample_indices: &[usize],
k: usize,
target_device: &Device,
) -> anyhow::Result<(Tensor, Tensor)> {
let n = sample_indices.len();
let mut idx_buf = vec![0u32; n * k];
let mut val_buf = vec![0.0f32; n * k];
idx_buf
.par_chunks_mut(k)
.zip(val_buf.par_chunks_mut(k))
.zip(sample_indices.par_iter())
.for_each(|((idx_chunk, val_chunk), &si)| {
let s = &samples[si];
let take = s.indices.len().min(k);
idx_chunk[..take].copy_from_slice(&s.indices[..take]);
val_chunk[..take].copy_from_slice(&s.values[..take]);
});
let indices =
Tensor::from_vec(idx_buf, (n, k), target_device)?.to_dtype(candle_core::DType::U32)?;
let values = Tensor::from_vec(val_buf, (n, k), target_device)?;
Ok((indices, values))
}
pub fn gather_per_feature_at_indices(
samples: &[IndexedSample],
sample_indices: &[usize],
per_feature: &[f32],
k: usize,
target_device: &Device,
) -> anyhow::Result<Tensor> {
pack_at_indices(
samples,
sample_indices,
k,
target_device,
|_, _, _, feat| per_feature[feat as usize],
)
}