Skip to main content

cubecl_ir/interfaces/
aliasing.rs

1use pliron::{basic_block::BasicBlock, value::DefiningEntity};
2
3use crate::prelude::*;
4
5#[op_interface]
6pub trait AliasingOp: OneResultInterface {
7    verify_op_succ!();
8    fn source_ptr(&self, ctx: &Context) -> Option<Value>;
9}
10
11pub trait PointerExt {
12    fn find_root_index(&self, ctx: &Context) -> usize;
13    fn get_root_defining_entity(&self, ctx: &Context) -> DefiningEntity;
14    fn get_root_defining_op(&self, ctx: &Context) -> Option<Ptr<Operation>> {
15        match self.get_root_defining_entity(ctx) {
16            DefiningEntity::Op(ptr) => Some(ptr),
17            DefiningEntity::Block(_) => None,
18        }
19    }
20    fn get_root_defining_block(&self, ctx: &Context) -> Option<Ptr<BasicBlock>> {
21        match self.get_root_defining_entity(ctx) {
22            DefiningEntity::Op(_) => None,
23            DefiningEntity::Block(ptr) => Some(ptr),
24        }
25    }
26    fn get_root_value(&self, ctx: &Context) -> Value {
27        let index = self.find_root_index(ctx);
28        match self.get_root_defining_entity(ctx) {
29            DefiningEntity::Op(op) => op.deref(ctx).get_result(index),
30            DefiningEntity::Block(block) => block.deref(ctx).get_argument(index),
31        }
32    }
33}
34
35impl PointerExt for Value {
36    fn get_root_defining_entity(&self, ctx: &Context) -> DefiningEntity {
37        match self.defining_entity() {
38            DefiningEntity::Op(op) => {
39                if let Some(aliasing) = op_cast::<dyn AliasingOp>(&*op.dyn_op(ctx))
40                    && let Some(source) = aliasing.source_ptr(ctx)
41                {
42                    source.get_root_defining_entity(ctx)
43                } else {
44                    DefiningEntity::Op(op)
45                }
46            }
47            block @ DefiningEntity::Block(_) => block,
48        }
49    }
50
51    fn find_root_index(&self, ctx: &Context) -> usize {
52        match self.defining_entity() {
53            DefiningEntity::Op(op)
54                if let Some(aliasing) = op_cast::<dyn AliasingOp>(&*op.dyn_op(ctx))
55                    && let Some(source) = aliasing.source_ptr(ctx) =>
56            {
57                source.find_root_index(ctx)
58            }
59            _ => self.find_index(ctx),
60        }
61    }
62}