use super::autotune::Optimization;
use crate::{
DType, Map, Set,
backend::DeviceInfo,
kernel::{BOp, Kernel, MemLayout, Op, OpId, UOp},
shape::Dim,
};
#[derive(Debug)]
pub struct Vectorize {
pub supported_lens: Vec<u8>,
pub vectorize_ops: bool,
}
impl Optimization for Vectorize {
fn nconfigs(&self) -> u64 {
1
}
fn apply(&self, kernel: &mut Kernel, _config: u64) {
kernel.vectorize_loads(&self.supported_lens);
kernel.vectorize_stores(&self.supported_lens);
if self.vectorize_ops {
kernel.vectorize_ops_forward(&self.supported_lens);
kernel.vectorize_ops_backward(&self.supported_lens);
}
}
}
#[derive(Debug)]
struct LoadInfo {
id: OpId,
index: OpId,
}
#[derive(Debug)]
struct StoreInfo {
id: OpId,
index: OpId,
x: OpId,
}
impl Kernel {
pub fn opt_vectorize(&self, dev_info: &DeviceInfo) -> Box<dyn Optimization> {
Box::new(Vectorize {
supported_lens: dev_info.supported_vec_lens.clone(),
vectorize_ops: !dev_info.supported_vec_lens.is_empty(),
})
}
pub fn vectorize_loads(&mut self, supported_lens: &[u8]) {
let mut op_id = self.head;
let mut loads: Vec<Map<OpId, Vec<LoadInfo>>> = Vec::new();
loads.push(Map::default());
while !op_id.is_null() {
match self.ops[op_id].op {
Op::Loop { .. } => {
loads.push(Map::default());
}
Op::Load { src, index, layout } => {
if layout == MemLayout::Scalar {
loads
.last_mut()
.unwrap()
.entry(src)
.and_modify(|e| e.push(LoadInfo { id: op_id, index }))
.or_insert_with(|| vec![LoadInfo { id: op_id, index }]);
}
}
Op::EndLoop => self.verify_and_apply_vectorization(&mut loads, supported_lens),
_ => {}
}
op_id = self.next_op(op_id);
}
self.verify_and_apply_vectorization(&mut loads, supported_lens);
}
fn verify_and_apply_vectorization(&mut self, loads: &mut Vec<Map<OpId, Vec<LoadInfo>>>, supported_lens: &[u8]) {
if let Some(loads) = loads.pop() {
for (src, mut loads) in loads {
if !supported_lens.contains(&(loads.len() as u8)) {
continue;
}
loads.sort_unstable_by_key(|x| self.get_strides(x.index).len());
let mut base_index = None;
let mut offset_order: Vec<Dim> = Vec::new();
let vec_len = loads.len() as Dim;
for (base_idx, (_, vl)) in self.get_strides(loads[0].index) {
if !(vl == vec_len || (base_idx.is_null() && vl == 0)) {
continue;
}
let mut offsets: Set<Dim> = (0..vec_len).collect();
offset_order.clear();
if loads[1..].iter().all(|x| {
let strides = self.get_strides(x.index);
if base_idx.is_null() {
strides.iter().any(|(&idx, (_, st))| {
let found = idx.is_null() && offsets.remove(st);
if found {
offset_order.push(*st);
}
found
})
} else {
strides.iter().any(|(&idx, (_, st))| idx == base_idx && *st == vec_len)
&& strides.iter().any(|(&idx, (_, st))| {
let found = idx.is_null() && offsets.remove(st);
if found {
offset_order.push(*st);
}
found
})
}
}) && offsets.remove(&0)
{
base_index = Some(base_idx);
break;
}
}
if base_index.is_some() {
let vload = self.insert_before(
loads[0].id,
Op::Load { src, index: loads[0].index, layout: MemLayout::Vector(vec_len as u16) },
);
self.ops[loads[0].id].op = Op::Index { vec: vload, idx: 0 };
for (load, &off) in loads[1..].iter().zip(&offset_order) {
self.ops[load.id].op = Op::Index { vec: vload, idx: off as usize };
}
}
}
}
}
pub fn vectorize_stores(&mut self, supported_lens: &[u8]) {
let mut op_id = self.head;
let mut stores: Vec<Map<OpId, Vec<StoreInfo>>> = Vec::new();
stores.push(Map::default());
while !op_id.is_null() {
match self.ops[op_id].op {
Op::Loop { .. } => {
stores.push(Map::default());
}
Op::Store { dst, src: x, index, layout } => {
if layout == MemLayout::Scalar {
stores
.last_mut()
.unwrap()
.entry(dst)
.and_modify(|e| e.push(StoreInfo { id: op_id, index, x }))
.or_insert_with(|| vec![StoreInfo { id: op_id, index, x }]);
}
}
Op::EndLoop => self.verify_and_apply_store_vectorization(&mut stores, supported_lens),
_ => {}
}
op_id = self.next_op(op_id);
}
self.verify_and_apply_store_vectorization(&mut stores, supported_lens);
}
fn verify_and_apply_store_vectorization(&mut self, stores: &mut Vec<Map<OpId, Vec<StoreInfo>>>, supported_lens: &[u8]) {
if let Some(stores) = stores.pop() {
for (dst, mut stores) in stores {
if !supported_lens.contains(&(stores.len() as u8)) {
continue;
}
stores.sort_unstable_by_key(|x| self.get_strides(x.index).len());
let mut base_index = None;
let mut offset_order: Vec<Dim> = Vec::new();
let vec_len = stores.len() as Dim;
for (base_idx, (_, vl)) in self.get_strides(stores[0].index) {
if !(vl == vec_len || (base_idx.is_null() && vl == 0)) {
continue;
}
let mut offsets: Set<Dim> = (0..vec_len).collect();
offset_order.clear();
if stores[1..].iter().all(|x| {
let strides = self.get_strides(x.index);
if base_idx.is_null() {
strides.iter().any(|(&idx, (_, st))| {
let found = idx.is_null() && offsets.remove(st);
if found {
offset_order.push(*st);
}
found
})
} else {
strides.iter().any(|(&idx, (_, st))| idx == base_idx && *st == vec_len)
&& strides.iter().any(|(&idx, (_, st))| {
let found = idx.is_null() && offsets.remove(st);
if found {
offset_order.push(*st);
}
found
})
}
}) && offsets.remove(&0)
{
base_index = Some(base_idx);
break;
}
}
if base_index.is_some() {
let mut ops = vec![OpId::NULL; vec_len as usize].into_boxed_slice();
ops[0] = stores[0].x;
for (store, &off) in stores[1..].iter().zip(&offset_order) {
ops[off as usize] = store.x;
}
let store_ids: Set<OpId> = stores.iter().map(|s| s.id).collect();
let mut last_id = stores[0].id;
let mut cur = stores[0].id;
while !cur.is_null() {
if store_ids.contains(&cur) {
last_id = cur;
}
cur = self.next_op(cur);
}
let vstore = self.insert_after(last_id, Op::Stack { ops });
self.insert_after(
vstore,
Op::Store { dst, src: vstore, index: stores[0].index, layout: MemLayout::Vector(vec_len as u16) },
);
for store in &stores {
self.remove_op(store.id);
}
}
}
}
}
fn op_is_before(&self, target: OpId, pos: OpId) -> bool {
let mut cur = self.head;
while !cur.is_null() && cur != pos {
if cur == target {
return true;
}
cur = self.next_op(cur);
}
false
}
fn last_in_kernel_order(&self, candidates: &[OpId]) -> OpId {
let set: Set<OpId> = candidates.iter().copied().collect();
let mut cur = self.tail;
while !cur.is_null() {
if set.contains(&cur) {
return cur;
}
cur = self.prev_op(cur);
}
candidates[0]
}
pub fn vectorize_ops_forward(&mut self, supported_lens: &[u8]) {
let mut supported: Vec<u8> = supported_lens.to_vec();
supported.sort_unstable_by(|a, b| b.cmp(a));
enum OpType {
Unary(UOp),
Cast(DType),
Bitcast(DType),
Binary(BOp, u8), }
#[allow(clippy::type_complexity)] let mut groups: Vec<(OpId, OpType, Vec<(OpId, usize)>)> = Vec::new();
loop {
groups.clear();
let mut op_id = self.head;
while !op_id.is_null() {
let info = match &self.ops[op_id].op {
Op::Unary { uop, x } => {
let (uop, x) = (*uop, *x);
match &self.ops[x].op {
Op::Index { vec, idx } => Some((*vec, OpType::Unary(uop), op_id, *idx)),
_ => None,
}
}
Op::Cast { dtype, x } => {
let (dtype, x) = (*dtype, *x);
match &self.ops[x].op {
Op::Index { vec, idx } => Some((*vec, OpType::Cast(dtype), op_id, *idx)),
_ => None,
}
}
Op::Bitcast { dtype, x } => {
let (dtype, x) = (*dtype, *x);
match &self.ops[x].op {
Op::Index { vec, idx } => Some((*vec, OpType::Bitcast(dtype), op_id, *idx)),
_ => None,
}
}
Op::Binary { bop, x, y } => {
let (bop, x, y) = (*bop, *x, *y);
if let Op::Index { vec, idx } = &self.ops[x].op {
Some((*vec, OpType::Binary(bop, 0), op_id, *idx))
} else if let Op::Index { vec, idx } = &self.ops[y].op {
Some((*vec, OpType::Binary(bop, 1), op_id, *idx))
} else {
None
}
}
_ => None,
};
let Some((source, op_type, consumer, _idx)) = info else {
op_id = self.next_op(op_id);
continue;
};
if let Some(g) = groups.iter_mut().find(|(s, t, _)| {
*s == source
&& match (t, &op_type) {
(OpType::Unary(a), OpType::Unary(b)) => a == b,
(OpType::Cast(a), OpType::Cast(b)) => a == b,
(OpType::Bitcast(a), OpType::Bitcast(b)) => a == b,
(OpType::Binary(a, ap), OpType::Binary(b, bp)) => a == b && ap == bp,
_ => false,
}
}) {
g.2.push((consumer, 0));
} else {
groups.push((source, op_type, vec![(consumer, 0)]));
}
op_id = self.next_op(op_id);
}
let mut indices: Vec<usize> = (0..groups.len()).filter(|&i| groups[i].2.len() >= 2).collect();
indices.sort_by(|&a, &b| groups[b].2.len().cmp(&groups[a].2.len()));
let mut applied = false;
for &idx in &indices {
let (_source, op_type, entries) = &groups[idx];
let n = entries.len();
let Some(vec_len) = supported.iter().copied().find(|&l| (l as usize) <= n) else {
continue;
};
let vec_len = vec_len as usize;
let selected: Vec<(OpId, usize)> = entries.iter().take(vec_len).copied().collect();
let first = selected[0].0;
if !matches!(op_type, OpType::Binary(_, _)) {
let lane_of = |consumer: OpId| match &self.ops[consumer].op {
Op::Unary { x, .. } | Op::Cast { x, .. } | Op::Bitcast { x, .. } => *x,
_ => unreachable!(),
};
if selected.iter().any(|&(c, _)| !self.op_is_before(lane_of(c), first)) {
continue;
}
}
if let OpType::Binary(_, devec_pos) = op_type {
let mut other_ops = Vec::with_capacity(vec_len);
for &(consumer, _) in &selected {
let (x, y) = match &self.ops[consumer].op {
Op::Binary { x, y, .. } => (*x, *y),
_ => unreachable!(),
};
let o = if *devec_pos == 0 { y } else { x };
other_ops.push(o);
}
if other_ops.iter().any(|&o| !self.op_is_before(o, first)) {
continue;
}
}
let vec_op_id = match op_type {
OpType::Unary(uop) => {
let ops: Box<[OpId]> = selected
.iter()
.map(|&(c, _)| match &self.ops[c].op {
Op::Unary { x, .. } => *x,
_ => unreachable!(),
})
.collect();
let vd = self.insert_before(first, Op::Stack { ops });
self.insert_before(first, Op::Unary { x: vd, uop: *uop })
}
OpType::Cast(dtype) => {
let ops: Box<[OpId]> = selected
.iter()
.map(|&(c, _)| match &self.ops[c].op {
Op::Cast { x, .. } => *x,
_ => unreachable!(),
})
.collect();
let vd = self.insert_before(first, Op::Stack { ops });
self.insert_before(first, Op::Cast { x: vd, dtype: *dtype })
}
OpType::Bitcast(dtype) => {
let ops: Box<[OpId]> = selected
.iter()
.map(|&(c, _)| match &self.ops[c].op {
Op::Bitcast { x, .. } => *x,
_ => unreachable!(),
})
.collect();
let vd = self.insert_before(first, Op::Stack { ops });
self.insert_before(first, Op::Bitcast { x: vd, dtype: *dtype })
}
OpType::Binary(bop, devec_pos) => {
let n = selected.len();
let mut devec_ops = Vec::with_capacity(n);
let mut other_ops = Vec::with_capacity(n);
for &(consumer, _) in &selected {
let (x, y) = match &self.ops[consumer].op {
Op::Binary { x, y, .. } => (*x, *y),
_ => unreachable!(),
};
if *devec_pos == 0 {
devec_ops.push(x);
other_ops.push(y);
} else {
devec_ops.push(y);
other_ops.push(x);
}
}
let vd = self.insert_before(first, Op::Stack { ops: devec_ops.into_boxed_slice() });
let vo = self.insert_before(first, Op::Stack { ops: other_ops.into_boxed_slice() });
let (vx, vy) = if *devec_pos == 0 { (vd, vo) } else { (vo, vd) };
self.insert_before(first, Op::Binary { x: vx, y: vy, bop: *bop })
}
};
for (i, &(consumer, _)) in selected.iter().enumerate() {
self.ops[consumer].op = Op::Index { vec: vec_op_id, idx: i };
}
applied = true;
break;
}
if !applied {
break;
}
}
}
pub fn vectorize_ops_backward(&mut self, supported_lens: &[u8]) {
let mut op_id = self.tail;
while !op_id.is_null() {
let ops = match &self.ops[op_id].op {
Op::Stack { ops } => ops.clone(),
_ => {
op_id = self.prev_op(op_id);
continue;
}
};
let n = ops.len();
if n < 2 || !supported_lens.contains(&(n as u8)) {
op_id = self.prev_op(op_id);
continue;
}
match &self.ops[ops[0]].op {
Op::Unary { uop, .. } => {
let uop = *uop;
let mut sources = Vec::with_capacity(n);
for &sub in ops.iter() {
match &self.ops[sub].op {
Op::Unary { x, uop: u } if *u == uop => sources.push(*x),
_ => {
sources.clear();
break;
}
}
}
if !sources.is_empty() && !sources.iter().any(|s| ops.contains(s)) {
let last = self.last_in_kernel_order(&sources);
let v_src = self.insert_after(last, Op::Stack { ops: sources.into_boxed_slice() });
let v_op = self.insert_after(v_src, Op::Unary { x: v_src, uop });
self.remap(op_id, v_op);
}
}
Op::Cast { dtype, .. } => {
let dtype = *dtype;
let mut sources = Vec::with_capacity(n);
for &sub in ops.iter() {
match &self.ops[sub].op {
Op::Cast { x, dtype: d } if *d == dtype => sources.push(*x),
_ => {
sources.clear();
break;
}
}
}
if !sources.is_empty() && !sources.iter().any(|s| ops.contains(s)) {
let last = self.last_in_kernel_order(&sources);
let v_src = self.insert_after(last, Op::Stack { ops: sources.into_boxed_slice() });
let v_op = self.insert_after(v_src, Op::Cast { x: v_src, dtype });
self.remap(op_id, v_op);
}
}
Op::Bitcast { dtype, .. } => {
let dtype = *dtype;
let mut sources = Vec::with_capacity(n);
for &sub in ops.iter() {
match &self.ops[sub].op {
Op::Bitcast { x, dtype: d } if *d == dtype => sources.push(*x),
_ => {
sources.clear();
break;
}
}
}
if !sources.is_empty() && !sources.iter().any(|s| ops.contains(s)) {
let last = self.last_in_kernel_order(&sources);
let v_src = self.insert_after(last, Op::Stack { ops: sources.into_boxed_slice() });
let v_op = self.insert_after(v_src, Op::Bitcast { x: v_src, dtype });
self.remap(op_id, v_op);
}
}
Op::Binary { bop, .. } => {
let bop = *bop;
let mut xs = Vec::with_capacity(n);
let mut ys = Vec::with_capacity(n);
for &sub in ops.iter() {
match &self.ops[sub].op {
Op::Binary { x, y, bop: b } if *b == bop => {
xs.push(*x);
ys.push(*y);
}
_ => {
xs.clear();
break;
}
}
}
if !xs.is_empty() && !xs.iter().any(|x| ops.contains(x)) && !ys.iter().any(|y| ops.contains(y)) {
let all: Vec<OpId> = xs.iter().chain(ys.iter()).copied().collect();
let last = self.last_in_kernel_order(&all);
let v_xs = self.insert_after(last, Op::Stack { ops: xs.into_boxed_slice() });
let v_ys = self.insert_after(v_xs, Op::Stack { ops: ys.into_boxed_slice() });
let v_op = self.insert_after(v_ys, Op::Binary { x: v_xs, y: v_ys, bop });
self.remap(op_id, v_op);
}
}
_ => {}
}
op_id = self.prev_op(op_id);
}
}
}
#[cfg(test)]
mod tests {
use crate::{
DType,
kernel::{BOp, Dev, Kernel, Op, UOp},
};
fn check_forward_result(k: &Kernel, c0: crate::kernel::OpId, c1: crate::kernel::OpId) {
assert!(matches!(k.ops[c0].op, Op::Index { .. }));
assert!(matches!(k.ops[c1].op, Op::Index { .. }));
let mut found_v = false;
let mut found_vop = false;
let mut op_id = k.head;
while !op_id.is_null() {
match &k.ops[op_id].op {
Op::Stack { ops } if ops.len() == 2 => found_v = true,
Op::Unary { uop: UOp::Cos, x } if matches!(k.ops[*x].op, Op::Stack { .. }) => found_vop = true,
_ => {}
}
op_id = k.next_op(op_id);
}
assert!(found_v, "Vectorize not found");
assert!(found_vop, "Vector cos not found");
}
#[test]
fn vectorize_ops_forward_2_lane() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 2);
let [s0, s1] = k.devectorize::<2>(vec_load);
let c0 = k.unary(s0, UOp::Cos);
let c1 = k.unary(s1, UOp::Cos);
k.store(dst, c0, g0);
let four = k.const_idx(4u32);
let idx_c1 = k.binary(g0, four, BOp::Add);
k.store(dst, c1, idx_c1);
k.vectorize_ops_forward(&[2]);
assert!(matches!(k.ops[s0].op, Op::Index { .. }));
assert!(matches!(k.ops[s1].op, Op::Index { .. }));
check_forward_result(&k, c0, c1);
}
#[test]
fn vectorize_ops_forward_4_lane() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 4);
let [s0, s1, s2, s3] = k.devectorize::<4>(vec_load);
let [c0, c1, c2, c3] = [s0, s1, s2, s3].map(|s| k.unary(s, UOp::Cos));
let c1i = k.const_idx(1u32);
let c2i = k.const_idx(2u32);
let c3i = k.const_idx(3u32);
let i1 = k.binary(g0, c1i, BOp::Add);
let i2 = k.binary(g0, c2i, BOp::Add);
let i3 = k.binary(g0, c3i, BOp::Add);
k.store(dst, c0, g0);
k.store(dst, c1, i1);
k.store(dst, c2, i2);
k.store(dst, c3, i3);
k.vectorize_ops_forward(&[2, 4]);
assert!(matches!(k.ops[s0].op, Op::Index { .. }));
assert!(matches!(k.ops[c0].op, Op::Index { .. }));
assert!(matches!(k.ops[c1].op, Op::Index { .. }));
assert!(matches!(k.ops[c2].op, Op::Index { .. }));
assert!(matches!(k.ops[c3].op, Op::Index { .. }));
let mut found_v = false;
let mut found_vop = false;
let mut found_cos_scalar = false;
let mut op_id = k.head;
while !op_id.is_null() {
match &k.ops[op_id].op {
Op::Stack { ops } if ops.len() == 4 => found_v = true,
Op::Unary { uop: UOp::Cos, x } if matches!(k.ops[*x].op, Op::Stack { .. }) => found_vop = true,
Op::Unary { uop: UOp::Cos, .. } => found_cos_scalar = true,
_ => {}
}
op_id = k.next_op(op_id);
}
assert!(found_v, "Vectorize(4) not found");
assert!(found_vop, "Vector cos(4) not found");
assert!(!found_cos_scalar, "No scalar cos should remain");
}
#[test]
fn vectorize_ops_forward_mixed_ops() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 4);
let [s0, s1, s2, s3] = k.devectorize::<4>(vec_load);
let c0 = k.unary(s0, UOp::Cos);
let c1 = k.unary(s1, UOp::Sin);
let c2 = k.unary(s2, UOp::Cos);
let c3 = k.unary(s3, UOp::Sin);
k.store(dst, c0, g0);
let c1i = k.const_idx(1u32);
let c2i = k.const_idx(2u32);
let c3i = k.const_idx(3u32);
let i1 = k.binary(g0, c1i, BOp::Add);
let i2 = k.binary(g0, c2i, BOp::Add);
let i3 = k.binary(g0, c3i, BOp::Add);
k.store(dst, c0, g0);
k.store(dst, c1, i1);
k.store(dst, c2, i2);
k.store(dst, c3, i3);
k.vectorize_ops_forward(&[2, 4]);
assert!(matches!(k.ops[c0].op, Op::Index { .. }));
assert!(matches!(k.ops[c1].op, Op::Index { .. }));
assert!(matches!(k.ops[c2].op, Op::Index { .. }));
assert!(matches!(k.ops[c3].op, Op::Index { .. }));
let mut v_count = 0;
let mut vop_count = 0;
let mut op_id = k.head;
while !op_id.is_null() {
match &k.ops[op_id].op {
Op::Stack { .. } => v_count += 1,
Op::Unary { uop: UOp::Cos, x } if matches!(k.ops[*x].op, Op::Stack { .. }) => vop_count += 1,
Op::Unary { uop: UOp::Sin, x } if matches!(k.ops[*x].op, Op::Stack { .. }) => vop_count += 1,
_ => {}
}
op_id = k.next_op(op_id);
}
assert_eq!(v_count, 2, "Two Vectorize ops expected");
assert_eq!(vop_count, 2, "Two vector Unary ops expected");
}
#[test]
fn vectorize_ops_forward_binary() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 2);
let [s0, s1] = k.devectorize::<2>(vec_load);
let c = k.const_val(1.0f32);
let r0 = k.binary(s0, c, BOp::Add);
let r1 = k.binary(s1, c, BOp::Add);
let four = k.const_idx(4u32);
let idx1 = k.binary(g0, four, BOp::Add);
k.store(dst, r0, g0);
k.store(dst, r1, idx1);
k.vectorize_ops_forward(&[2]);
assert!(matches!(k.ops[s0].op, Op::Index { .. }));
assert!(matches!(k.ops[s1].op, Op::Index { .. }));
assert!(matches!(k.ops[r0].op, Op::Index { .. }));
assert!(matches!(k.ops[r1].op, Op::Index { .. }));
let mut found_vec_bin = false;
let mut op_id = k.head;
while !op_id.is_null() {
if let Op::Binary { bop: BOp::Add, x, y } = &k.ops[op_id].op
&& matches!(k.ops[*x].op, Op::Stack { .. })
&& matches!(k.ops[*y].op, Op::Stack { .. })
{
found_vec_bin = true;
}
op_id = k.next_op(op_id);
}
assert!(found_vec_bin, "Vector Binary(Add) op not found");
}
#[test]
fn vectorize_ops_forward_binary_y_pos() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 2);
let [s0, s1] = k.devectorize::<2>(vec_load);
let c = k.const_val(1.0f32);
let r0 = k.binary(c, s0, BOp::Add); let r1 = k.binary(c, s1, BOp::Add); let four = k.const_idx(4u32);
let idx1 = k.binary(g0, four, BOp::Add);
k.store(dst, r0, g0);
k.store(dst, r1, idx1);
k.vectorize_ops_forward(&[2]);
assert!(matches!(k.ops[r0].op, Op::Index { .. }));
assert!(matches!(k.ops[r1].op, Op::Index { .. }));
}
#[test]
fn vectorize_ops_and_constfold_clears_vectorize_devectorize() {
let mut k = Kernel::from_device_id(Dev::Auto, None);
let src = k.param(DType::F32);
let dst = k.param(DType::F32);
let g0_len = k.const_idx(4);
let g0 = k.group_range(0, g0_len);
let two = k.const_idx(2u32);
let offset = k.binary(g0, two, BOp::BitShiftLeft);
let vec_load = k.load_vector(src, offset, 4);
let [s0, s1, s2, s3] = k.devectorize::<4>(vec_load);
let c0 = k.unary(s0, UOp::Cos);
let c1 = k.unary(s1, UOp::Cos);
let c2 = k.unary(s2, UOp::Cos);
let c3 = k.unary(s3, UOp::Cos);
let vec = k.stack(&[c0, c1, c2, c3]);
k.store_vector(dst, vec, offset, 4);
k.vectorize_ops_backward(&[4]);
k.constant_folding();
k.dead_code_elimination();
let mut op_id = k.head;
while !op_id.is_null() {
match k.ops[op_id].op {
Op::Stack { .. } | Op::Index { .. } => {
panic!("Found Vectorize/Devectorize op at {op_id} after passes");
}
_ => {}
}
op_id = k.next_op(op_id);
}
}
}