Skip to main content

cubecl_opt/analyses/alias_analysis/
mod.rs

1use core::{
2    cell::RefMut,
3    ops::{BitAndAssign, BitOrAssign},
4};
5
6use alloc::{boxed::Box, vec::Vec};
7use cubecl_ir::{interfaces::MemoryEffect, prelude::*};
8use pliron::{context::Context, pass::AnalysisManager, value::Value};
9
10use crate::analyses::{
11    alias_analysis::{address_space::AddressSpaceAA, root_alloc::RootAllocAA},
12    memory_ssa::MemorySSA,
13};
14
15pub mod address_space;
16pub mod root_alloc;
17
18#[derive(PartialEq, Eq, PartialOrd, Ord, Clone, Copy)]
19pub enum ModRefResult {
20    NoModRef,
21    Ref,
22    Mod,
23    ModRef,
24}
25
26impl ModRefResult {
27    pub fn join(&self, other: ModRefResult) -> ModRefResult {
28        match (*self, other) {
29            (ModRefResult::NoModRef, other) | (other, ModRefResult::NoModRef) => other,
30            (_, ModRefResult::ModRef) | (ModRefResult::ModRef, _) => ModRefResult::ModRef,
31            (ModRefResult::Ref, ModRefResult::Ref) => ModRefResult::Ref,
32            (ModRefResult::Mod, ModRefResult::Mod) => ModRefResult::Mod,
33            (ModRefResult::Ref, ModRefResult::Mod) | (ModRefResult::Mod, ModRefResult::Ref) => {
34                ModRefResult::ModRef
35            }
36        }
37    }
38
39    pub fn meet(&self, other: ModRefResult) -> ModRefResult {
40        match (*self, other) {
41            (ModRefResult::NoModRef, _) | (_, ModRefResult::NoModRef) => ModRefResult::NoModRef,
42            (ModRefResult::ModRef, other) | (other, ModRefResult::ModRef) => other,
43            (ModRefResult::Ref, ModRefResult::Ref) => ModRefResult::Ref,
44            (ModRefResult::Mod, ModRefResult::Mod) => ModRefResult::Mod,
45            (ModRefResult::Mod, ModRefResult::Ref) | (ModRefResult::Ref, ModRefResult::Mod) => {
46                ModRefResult::NoModRef
47            }
48        }
49    }
50
51    pub fn contains_mod(&self) -> bool {
52        match self {
53            ModRefResult::NoModRef | ModRefResult::Ref => false,
54            ModRefResult::Mod | ModRefResult::ModRef => true,
55        }
56    }
57}
58
59impl BitOrAssign for ModRefResult {
60    fn bitor_assign(&mut self, rhs: Self) {
61        *self = self.join(rhs);
62    }
63}
64
65impl BitAndAssign for ModRefResult {
66    fn bitand_assign(&mut self, rhs: Self) {
67        *self = self.meet(rhs);
68    }
69}
70
71#[derive(Clone, Copy, PartialEq, Eq)]
72pub enum AliasResult {
73    NoAlias,
74    MayAlias,
75    PartialAlias,
76    MustAlias,
77}
78
79impl AliasResult {
80    pub fn join(&self, other: AliasResult) -> AliasResult {
81        if *self == other {
82            return other;
83        }
84        match (self, other) {
85            (AliasResult::PartialAlias, AliasResult::MustAlias)
86            | (AliasResult::MustAlias, AliasResult::PartialAlias) => AliasResult::PartialAlias,
87            _ => AliasResult::MayAlias,
88        }
89    }
90}
91
92pub trait AliasAnalysis {
93    fn alias(&self, ctx: &Context, lhs: Value, rhs: Value) -> AliasResult;
94    fn mod_ref(&self, ctx: &Context, location: &MemoryEffect, rhs: &MemoryEffect) -> ModRefResult;
95}
96
97#[derive(Default)]
98pub struct AliasAnalysisStack {
99    analyses: Vec<Box<dyn AliasAnalysis>>,
100}
101
102impl AliasAnalysisStack {
103    pub fn add_analysis(&mut self, analysis: impl AliasAnalysis + 'static) {
104        self.analyses.push(Box::new(analysis));
105    }
106
107    pub fn alias(&self, ctx: &Context, lhs: Value, rhs: Value) -> AliasResult {
108        let mut result = AliasResult::MayAlias;
109        for analysis in self.analyses.iter() {
110            result = analysis.alias(ctx, lhs, rhs);
111            if result != AliasResult::MayAlias {
112                return result;
113            }
114        }
115        result
116    }
117
118    pub fn mod_ref<'a>(
119        &self,
120        ctx: &Context,
121        location: impl IntoIterator<Item = &'a MemoryEffect> + Clone,
122        rhs: impl IntoIterator<Item = &'a MemoryEffect> + Clone,
123    ) -> ModRefResult {
124        // Within an effect pair: meet AA results, since each analysis refines the
125        // conservative answer.
126        //
127        // Across effect pairs: join results, since either pair may account for the
128        // interaction between the two effect sets.
129        let mut result = ModRefResult::NoModRef;
130        for lhs_effect in location.clone() {
131            for rhs_effect in rhs.clone() {
132                let mut effect_res = ModRefResult::ModRef;
133                for analysis in self.analyses.iter() {
134                    effect_res &= analysis.mod_ref(ctx, lhs_effect, rhs_effect);
135                }
136                result |= effect_res;
137            }
138        }
139
140        result
141    }
142}
143
144pub(crate) fn effect_mod_ref(effect: &MemoryEffect) -> ModRefResult {
145    match effect {
146        MemoryEffect::Read(_) | MemoryEffect::ReadAllInSpace(_) | MemoryEffect::ReadAll => {
147            ModRefResult::Ref
148        }
149        MemoryEffect::Write(_) | MemoryEffect::WriteAllInSpace(_) | MemoryEffect::WriteAll => {
150            ModRefResult::Mod
151        }
152        MemoryEffect::Opaque => ModRefResult::ModRef,
153    }
154}
155
156pub fn default_stack() -> AliasAnalysisStack {
157    let mut stack = AliasAnalysisStack::default();
158    stack.add_analysis(AddressSpaceAA);
159    stack.add_analysis(RootAllocAA);
160    stack
161}
162
163pub fn default_memory_ssa<'a>(
164    ctx: &Context,
165    op: Ptr<Operation>,
166    analyses: &'a mut AnalysisManager,
167) -> Result<RefMut<'a, MemorySSA>> {
168    let mut memory_ssa = analyses.get_analysis_mut::<MemorySSA>(op, ctx)?;
169    memory_ssa.set_alias_analysis_stack(default_stack());
170    Ok(memory_ssa)
171}