use crate::algebra::group_key::{cell_image, for_cell_width, GroupKey, KeyCells};
use crate::algebra::reindex::{FoldCols, ReindexPacker};
use crate::repr::{Batch, MemBatch};
use crate::schema::Slot;
use crate::schema::{worker_for_key, worker_for_pk_bytes};
use crate::schema::{Placement, SchemaDescriptor};
use gnitz_wire::zip_cells;
pub fn op_worker_filter(batch: &Batch, slot: Slot) -> Batch {
ScatterPlan::native(Placement::full_pk(batch.schema())).share(batch, slot)
}
pub struct ScatterPlan(Option<GroupKey>);
impl ScatterPlan {
pub fn group(schema: &SchemaDescriptor, cols: &[u32]) -> Result<Self, String> {
Ok(ScatterPlan(Some(GroupKey::new(schema, cols)?)))
}
pub fn join(schema: &SchemaDescriptor, slots: &[gnitz_wire::ReindexSlot]) -> Result<Self, String> {
Ok(ScatterPlan(Some(GroupKey::packed(ReindexPacker::new(schema, slots)?))))
}
pub fn native(placement: Placement) -> Self {
match placement {
Placement::Replicated => Self::broadcast(),
Placement::Keyed { dist_stride } => ScatterPlan(Some(GroupKey::PkRange { at: 0, n: dist_stride as usize })),
Placement::Local => panic!("ScatterPlan::native: a Local relation's rows have no key owner"),
}
}
pub fn broadcast() -> Self {
ScatterPlan(None)
}
pub fn routes_to_native_owner(&self, placement: Placement) -> bool {
matches!((&self.0, placement),
(Some(GroupKey::PkRange { at: 0, n }), Placement::Keyed { dist_stride }) if *n == dist_stride as usize)
}
pub fn route<'a>(&self, batch: &Batch, out: &'a mut Vec<Vec<u32>>, num_workers: usize) -> &'a [Vec<u32>] {
let lists = match self.0 {
None => 1,
Some(_) => num_workers,
};
if out.len() < lists {
out.resize_with(lists, Vec::new);
}
let slots = &mut out[..lists];
slots.iter_mut().for_each(Vec::clear);
self.route_into(&batch.as_mem_batch(), num_workers, &mut EveryWorker(slots));
slots
}
pub fn share(&self, batch: &Batch, slot: Slot) -> Batch {
let nw = slot.of as usize;
let mut share = Share {
rank: match self.0 {
None => 0,
Some(_) => slot.rank as usize,
},
rows: Vec::with_capacity(batch.count / nw + 1),
};
self.route_into(&batch.as_mem_batch(), nw, &mut share);
batch.ascending_subset(&share.rows)
}
fn route_into(&self, mb: &MemBatch, nw: usize, sink: &mut impl RowSink) {
match self.0 {
_ if nw == 1 => route_rows(mb, sink, ListZero),
Some(ref key) => match key.cells(mb) {
Some(cells) => for_cell_width!(cells.width, |W| match cells.opk {
true => route_cells::<W, true>(&cells, mb.weight(), sink, nw),
false => route_cells::<W, false>(&cells, mb.weight(), sink, nw),
}),
None => match *key {
GroupKey::PkRange { at, n } => route_rows(mb, sink, PkRangeW { at, n, nw }),
GroupKey::Packed(ref packer) => route_rows_packed(mb, sink, packer, nw),
GroupKey::Fold(ref fold) => route_rows(mb, sink, FoldW { fold, nw }),
GroupKey::Image(_) => unreachable!("an image key is one cell"),
},
},
None => route_rows(mb, sink, ListZero),
}
}
}
trait RowSink {
fn put(&mut self, worker: usize, row: u32);
}
struct EveryWorker<'a>(&'a mut [Vec<u32>]);
impl RowSink for EveryWorker<'_> {
#[inline(always)]
fn put(&mut self, worker: usize, row: u32) {
self.0[worker].push(row);
}
}
struct Share {
rank: usize,
rows: Vec<u32>,
}
impl RowSink for Share {
#[inline(always)]
fn put(&mut self, worker: usize, row: u32) {
if worker == self.rank {
self.rows.push(row);
}
}
}
trait RowWorker {
fn worker(&self, mb: &MemBatch, row: usize) -> usize;
}
struct PkRangeW {
at: usize,
n: usize,
nw: usize,
}
impl RowWorker for PkRangeW {
#[inline(always)]
fn worker(&self, mb: &MemBatch, row: usize) -> usize {
worker_for_pk_bytes(mb.get_pk_range(row, self.at, self.n), self.nw)
}
}
struct ListZero;
impl RowWorker for ListZero {
#[inline(always)]
fn worker(&self, _: &MemBatch, _: usize) -> usize {
0
}
}
struct FoldW<'a> {
fold: &'a FoldCols,
nw: usize,
}
impl RowWorker for FoldW<'_> {
#[inline(always)]
fn worker(&self, mb: &MemBatch, row: usize) -> usize {
worker_for_key(self.fold.key_row(mb, row, mb.get_null_word(row)), self.nw)
}
}
#[inline(never)]
fn route_rows<W: RowWorker>(mb: &MemBatch, sink: &mut impl RowSink, w: W) {
for (row, weight) in mb.weight().as_chunks::<8>().0.iter().enumerate() {
if *weight != [0; 8] {
sink.put(w.worker(mb, row), row as u32);
}
}
}
#[inline(never)]
fn route_cells<const W: usize, const OPK: bool>(cells: &KeyCells, weights: &[u8], sink: &mut impl RowSink, nw: usize) {
let rows = weights.as_chunks::<8>().0.iter().enumerate();
zip_cells::<W, _>(cells.region, cells.stride, cells.off, rows, |cell, (row, weight)| {
if *weight != [0; 8] {
sink.put(worker_for_key(cell_image::<W, OPK>(cell, cells.signed), nw), row as u32);
}
});
}
#[inline(never)]
fn route_rows_packed(mb: &MemBatch, sink: &mut impl RowSink, packer: &ReindexPacker, nw: usize) {
packer.for_each_key(mb, packer.out_stride, |row, key| {
if mb.get_weight(row) != 0 {
sink.put(worker_for_pk_bytes(key, nw), row as u32);
}
});
}
#[cfg(test)]
#[path = "tests/exchange.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/exchange.rs"]
mod bench;