cubecl_cpp/shared/
unroll.rs1use 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#[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 *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}