Skip to main content

cubecl_cpp/shared/
unroll.rs

1use cubecl_core::ir::{
2    NamedRewrite,
3    dialect::vector::{CompositeConstructOp, CompositeExtractOp},
4    interfaces::{MaterializableOp, TypedExt},
5    prelude::*,
6    rewrite::MatchRewritePass,
7    types::VectorType,
8};
9
10#[op_interface]
11pub trait UnrollingOp: MaterializableOp + OneResultInterface {
12    verify_op_succ!();
13}
14
15macro_rules! unrolling {
16    ($ty: ty) => {
17        #[pliron::derive::op_interface_impl]
18        impl crate::shared::unroll::UnrollingOp for $ty {}
19    };
20}
21pub(crate) use unrolling;
22
23pub type CppUnrollPass = MatchRewritePass<CppUnroll>;
24
25// Different implementation because the semantics are a bit different than the full unroll pass.
26// It's also much simpler.
27#[derive(Default, Clone, NamedRewrite)]
28pub struct CppUnroll;
29
30impl MatchRewrite for CppUnroll {
31    fn r#match(&mut self, ctx: &Context, op: Ptr<Operation>) -> bool {
32        op.impls::<dyn UnrollingOp>(ctx) && op.result(ctx).is_vector(ctx)
33    }
34
35    fn rewrite(
36        &mut self,
37        ctx: &mut Context,
38        rewriter: &mut MatchRewriter,
39        op: Ptr<Operation>,
40    ) -> Result<()> {
41        let dyn_op = op.dyn_op(ctx);
42        let unroll_op = op_cast::<dyn UnrollingOp>(&*dyn_op).unwrap();
43        let opds = op.operands(ctx);
44        let attributes = op.deref(ctx).attributes.clone();
45        let res = unroll_op.get_result(ctx);
46
47        let vec_ty = {
48            let ty = res.get_type(ctx).deref(ctx);
49            *ty.downcast_ref::<VectorType>().unwrap()
50        };
51
52        let extract = |ctx: &mut Context, opd: &Value, i: usize| {
53            if opd.vector_size(ctx) == 1 {
54                // Scalar arg for things like lane index in plane ops. SameOperandTypes should
55                // validate other args for equality so we don't get implicit broadcasts
56                *opd
57            } else {
58                let extract = CompositeExtractOp::new(ctx, *opd, i);
59                extract.get_operation().insert_before(ctx, op);
60                extract.get_result(ctx)
61            }
62        };
63        let run_one = |ctx: &mut Context, i: usize| {
64            let opds = opds.iter().map(|opd| extract(ctx, opd, i)).collect();
65            let attrs = attributes.clone();
66            let new_op = unroll_op.materialize(ctx, vec![vec_ty.inner], opds, attrs);
67            new_op.insert_before(ctx, op);
68            new_op.result(ctx)
69        };
70
71        let new_values = (0..vec_ty.vectorization).map(|i| run_one(ctx, i)).collect();
72        let new_vec = CompositeConstructOp::new(ctx, res.get_type(ctx), new_values);
73        new_vec.get_operation().insert_before(ctx, op);
74        rewriter.replace_operation(ctx, op, new_vec.get_operation());
75
76        Ok(())
77    }
78}