cubecl_cpp/metal/
address_space.rs1use 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 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 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}