Skip to main content

cubecl_cpp/metal/
address_space.rs

1use cubecl_core::prelude::Visibility;
2
3use crate::{
4    Dialect,
5    shared::{Component, Item, KernelArg, PointerClass, Value},
6};
7
8use super::BufferAttribute;
9use std::fmt::Display;
10
11#[derive(Debug, PartialEq, Eq, Clone, Copy)]
12pub enum AddressSpace {
13    Constant,
14    ConstDevice,
15    Device,
16    Thread,
17    ThreadGroup,
18    None,
19}
20
21impl AddressSpace {
22    pub fn attribute(&self) -> BufferAttribute {
23        match self {
24            AddressSpace::Constant | AddressSpace::ConstDevice | AddressSpace::Device => {
25                BufferAttribute::Buffer
26            }
27            AddressSpace::ThreadGroup => BufferAttribute::ThreadGroup,
28            _ => BufferAttribute::None,
29        }
30    }
31}
32
33impl Display for AddressSpace {
34    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
35        match self {
36            AddressSpace::Constant => f.write_str("constant"),
37            AddressSpace::ConstDevice => f.write_str("const device"),
38            AddressSpace::Device => f.write_str("device"),
39            AddressSpace::ThreadGroup => f.write_str("threadgroup"),
40            AddressSpace::Thread => f.write_str("thread"),
41            AddressSpace::None => Ok(()),
42        }
43    }
44}
45
46impl From<AddressSpace> for Visibility {
47    fn from(val: AddressSpace) -> Self {
48        match val {
49            AddressSpace::Constant => Visibility::Read,
50            _ => Visibility::ReadWrite,
51        }
52    }
53}
54
55impl<D: Dialect> From<&KernelArg<D>> for AddressSpace {
56    fn from(value: &KernelArg<D>) -> Self {
57        // No atomic guard needed: this only feeds `.attribute()`, which maps every
58        // device space to `BufferAttribute::Buffer`. The atomic mutability constraint
59        // is enforced in `From<&Value>` and `compile_item` where it's observable.
60        value.vis.into()
61    }
62}
63
64impl From<Visibility> for AddressSpace {
65    fn from(value: Visibility) -> Self {
66        match value {
67            Visibility::Read => AddressSpace::ConstDevice,
68            Visibility::ReadWrite => AddressSpace::Device,
69            Visibility::Uniform => AddressSpace::Constant,
70        }
71    }
72}
73
74impl<D: Dialect> From<&Value<D>> for AddressSpace {
75    fn from(value: &Value<D>) -> Self {
76        if let Item::Pointer(inner, class) = value.item() {
77            // Atomics always need mutable (device) access, even on read-only
78            // bindings, because MSL forbids `const`-qualified `atomic<T>` pointers.
79            if matches!(inner.value_ty(), Item::Atomic(_))
80                && let PointerClass::Global(_) = class
81            {
82                return AddressSpace::Device;
83            }
84            return match class {
85                PointerClass::Global(visibility) => visibility.into(),
86                PointerClass::Shared => AddressSpace::ThreadGroup,
87                PointerClass::Local => AddressSpace::Thread,
88            };
89        }
90        AddressSpace::Thread
91    }
92}