Skip to main content

cubecl_cpp/cuda/
packed_ops.rs

1use cubecl_core::ir::{
2    NamedRewrite,
3    dialect::general::ReinterpretCastOp,
4    interfaces::{TypedExt, ValueExt},
5    prelude::*,
6    rewrite::MatchRewritePass,
7};
8
9#[op_interface]
10pub trait PackableOp: OneResultInterface {
11    verify_op_succ!();
12    fn should_pack(&self, ctx: &Context) -> bool;
13}
14
15macro_rules! packable {
16    ($ty: ty) => {
17        #[op_interface_impl]
18        impl crate::cuda::packed_ops::PackableOp for $ty {
19            fn should_pack(&self, ctx: &pliron::context::Context) -> bool {
20                use crate::shared::ty::TypedExtCPP;
21                self.get_result(ctx).can_pack(ctx)
22            }
23        }
24    };
25}
26pub(crate) use packable;
27use pliron::irbuild::inserter::Inserter;
28
29use crate::shared::ty::TypedExtCPP;
30
31pub type PackOpsPass = MatchRewritePass<PackOps>;
32
33#[derive(Default, Clone, NamedRewrite)]
34pub struct PackOps;
35
36impl MatchRewrite for PackOps {
37    fn r#match(&mut self, ctx: &Context, op: Ptr<Operation>) -> bool {
38        op_cast::<dyn PackableOp>(&*op.dyn_op(ctx)).is_some_and(|it| it.should_pack(ctx))
39    }
40
41    fn rewrite(
42        &mut self,
43        ctx: &mut Context,
44        rewriter: &mut MatchRewriter,
45        op: Ptr<Operation>,
46    ) -> Result<()> {
47        let dyn_op = op.dyn_op(ctx);
48        let unroll_op = op_cast::<dyn PackableOp>(&*dyn_op).unwrap();
49        let res = unroll_op.get_result(ctx);
50        let res_ty = res.get_type(ctx);
51
52        for r#use in op.operands_as_uses(ctx) {
53            let value = r#use.get_def(ctx);
54            // Skip scalar arg in plane ops
55            if value.vector_size(ctx) > 1 {
56                let reinterpret = ReinterpretCastOp::new(ctx, value.packed_type(ctx), value);
57                value.replace_use_with(ctx, r#use, &reinterpret.get_result(ctx));
58                reinterpret.get_operation().insert_before(ctx, op);
59            }
60        }
61
62        rewriter.set_value_type(ctx, res, res_ty.packed_type(ctx));
63        rewriter.set_insertion_point_after_operation(op);
64        let reinterpret_res = ReinterpretCastOp::new(ctx, res_ty, res);
65        rewriter.insert_op(ctx, &reinterpret_res);
66        let new_res = reinterpret_res.get_result(ctx);
67        res.replace_all_uses_except_with(ctx, reinterpret_res.input_as_use(ctx), &new_res);
68
69        Ok(())
70    }
71}