Skip to main content

cubecl_cpp/metal/
address_space.rs

1use cubecl_core::prelude::Visibility;
2
3use super::BufferAttribute;
4use std::fmt::Display;
5
6#[derive(Debug, PartialEq, Eq, Clone, Copy)]
7pub enum AddressSpace {
8    Constant,
9    ConstDevice,
10    Device,
11    Thread,
12    ThreadGroup,
13    None,
14}
15
16impl AddressSpace {
17    pub fn attribute(&self) -> BufferAttribute {
18        match self {
19            AddressSpace::Constant | AddressSpace::ConstDevice | AddressSpace::Device => {
20                BufferAttribute::Buffer
21            }
22            AddressSpace::ThreadGroup => BufferAttribute::ThreadGroup,
23            _ => BufferAttribute::None,
24        }
25    }
26}
27
28impl Display for AddressSpace {
29    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
30        match self {
31            AddressSpace::Constant => f.write_str("constant"),
32            AddressSpace::ConstDevice => f.write_str("const device"),
33            AddressSpace::Device => f.write_str("device"),
34            AddressSpace::ThreadGroup => f.write_str("threadgroup"),
35            AddressSpace::Thread => f.write_str("thread"),
36            AddressSpace::None => Ok(()),
37        }
38    }
39}
40
41impl From<AddressSpace> for Visibility {
42    fn from(val: AddressSpace) -> Self {
43        match val {
44            AddressSpace::Constant => Visibility::Read,
45            _ => Visibility::ReadWrite,
46        }
47    }
48}
49
50impl From<Visibility> for AddressSpace {
51    fn from(value: Visibility) -> Self {
52        match value {
53            Visibility::Read => AddressSpace::ConstDevice,
54            Visibility::ReadWrite => AddressSpace::Device,
55            Visibility::Uniform => AddressSpace::Constant,
56        }
57    }
58}