use std::collections::HashSet;
use furiosa_mapping::*;
pub(crate) fn transpose_broadcast<Src: M, Dst: M>(allow_broadcast: bool) -> Mapping {
let src_mapping = Src::to_value();
let dst_mapping = Dst::to_value();
let broadcast = dst_mapping.carve(&src_mapping);
if !allow_broadcast {
assert!(broadcast.is_padding());
}
broadcast
}
pub(crate) fn broadcast_axes(src: &Mapping, dst: &Mapping) -> Mapping {
let src_idents: HashSet<Ident> = src.idents().into_iter().collect();
let axes: Vec<Term> = dst
.axes()
.into_iter()
.filter(|term| !src_idents.contains(&term.symbol))
.map(AxisTerm::to_term)
.collect();
Mapping::from_terms(axes)
}
pub(crate) fn scatter_params(src: &Mapping, dst: &Mapping, key: &Mapping) -> (Mapping, AxisTerm) {
let payload = src.carve(key);
let dst_term = dst
.carve(&payload)
.axes()
.into_iter()
.next()
.expect("scatter dst residue has no live target axis");
(payload, dst_term)
}
pub(crate) struct GatherParams {
pub payload: Mapping,
pub src_term: AxisTerm,
}
pub(crate) fn gather_params(src: &Mapping, dst: &Mapping, idx: &Mapping) -> GatherParams {
let payload = dst.carve(idx);
let _ = dst.carve(&payload);
let src_term = src
.carve(&payload)
.axes()
.into_iter()
.next()
.expect("gather src residue has no live target axis");
GatherParams { payload, src_term }
}
#[cfg(test)]
mod tests {
use furiosa_mapping::*;
use super::broadcast_axes;
axes![A = 4, B = 2, C = 8];
#[test]
fn broadcast_axes_matches_carve() {
let cases: [(Mapping, Mapping); 4] = [
(<m![A, C]>::to_value(), <m![A, B, C]>::to_value()), (<m![A]>::to_value(), <m![A, B]>::to_value()),
(<m![B]>::to_value(), <m![A, B, C]>::to_value()), (<m![A, B, C]>::to_value(), <m![A, B, C]>::to_value()), ];
let sorted_axes = |m: &Mapping| {
let mut a = m.axes();
a.sort();
a
};
for (src, dst) in cases {
assert_eq!(
sorted_axes(&broadcast_axes(&src, &dst)),
sorted_axes(&dst.carve(&src)),
"broadcast_axes(src={src:?}, dst={dst:?})"
);
}
}
}