cubecl_cpp/cuda/
packed_ops.rs1use 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 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}