use cubecl_core::{
ir::{ConstantValue, Id},
prelude::Visibility,
};
use cubecl_ir::Intern;
use std::fmt::Display;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Value {
Constant(ConstantValue, Item),
Value { id: Id, item: Item },
}
#[derive(Debug, Clone, PartialEq)]
pub enum Builtin {
Id,
LocalInvocationIndex,
LocalInvocationIdX,
LocalInvocationIdY,
LocalInvocationIdZ,
WorkgroupId,
WorkgroupIdX,
WorkgroupIdY,
WorkgroupIdZ,
GlobalInvocationIdX,
GlobalInvocationIdY,
GlobalInvocationIdZ,
WorkgroupSize,
WorkgroupSizeX,
WorkgroupSizeY,
WorkgroupSizeZ,
NumWorkgroups,
NumWorkgroupsX,
NumWorkgroupsY,
NumWorkgroupsZ,
SubgroupSize,
SubgroupId,
SubgroupInvocationId,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Copy)]
pub enum Elem {
F16,
F32,
F64,
I32,
I64,
U32,
U64,
Bool,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq, Copy)]
pub enum Item {
Vector(Elem, usize),
Scalar(Elem),
Atomic(Intern<Item>),
Pointer(Intern<Item>, PointerClass),
Array(Intern<Item>, usize),
DynamicArray(Intern<Item>),
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Copy)]
pub enum PointerClass {
Global(Visibility),
Shared,
Local,
}
#[derive(Debug, Clone)]
pub struct IndexedValue {
val: Value,
index: usize,
}
impl Value {
pub fn index(&self, index: usize) -> IndexedValue {
IndexedValue {
val: self.clone(),
index,
}
}
pub fn is_ptr(&self) -> bool {
self.item().is_ptr()
}
pub fn item(&self) -> Item {
match self {
Self::Value { item, .. } => *item,
Self::Constant(_, item) => *item,
}
}
pub fn elem(&self) -> Elem {
*self.item().elem()
}
pub fn fmt_cast_to(&self, mut item: Item) -> String {
while let Item::Pointer(inner, _) = item {
item = *inner;
}
if self.item() == item || self.is_ptr() {
return format!("{self}");
}
let from = self.item();
let from_elem = *from.elem();
let to_elem = *item.elem();
let is_64bit = matches!(from_elem, Elem::I64 | Elem::U64);
let is_32bit_target = matches!(to_elem, Elem::I32 | Elem::U32);
if is_64bit && is_32bit_target {
let bitcast_elem = if matches!(to_elem, Elem::U32) {
Elem::U64
} else {
Elem::I64
};
if matches!(from, Item::Scalar(_)) {
let scalar_cast = format!("{to_elem}(bitcast<{bitcast_elem}>({self}))");
if matches!(item, Item::Scalar(_)) {
return scalar_cast;
}
return format!("{item}({scalar_cast})");
}
let bitcast_item = from.with_elem(bitcast_elem);
return format!("{item}(bitcast<{bitcast_item}>({self}))");
}
if from_elem == Elem::Bool && to_elem == Elem::F16 {
let f32_item = from.with_elem(Elem::F32);
return format!("{item}({f32_item}({self}))");
}
match (from, item) {
(Item::Scalar(_), Item::Scalar(_)) => format!("{item}({self})"),
(_, Item::Scalar(_)) => format!("{item}({self}.x)"),
(Item::Scalar(_), _) if from_elem != to_elem => format!("{item}({to_elem}({self}))"),
_ => format!("{item}({self})"),
}
}
}
impl Item {
pub fn intern(self) -> Intern<Self> {
Intern::new(self)
}
pub fn elem(&self) -> &Elem {
match self {
Item::Scalar(e) => e,
Item::Vector(elem, _) => elem,
Item::Atomic(inner) => inner.elem(),
Item::Pointer(inner, _) => inner.elem(),
Item::Array(inner, _) => inner.elem(),
Item::DynamicArray(inner) => inner.elem(),
}
}
pub fn unwrap_ptr(&self) -> Item {
match self {
Item::Pointer(inner, _) => **inner,
other => *other,
}
}
pub fn size(&self) -> usize {
match self {
Item::Scalar(e) => e.size(),
Item::Vector(elem, vector_size) => elem.size() * *vector_size,
Item::Atomic(inner) => inner.size(),
Item::Array(inner, length) => inner.size() * *length,
Item::DynamicArray(inner) => inner.size(),
Item::Pointer(..) => size_of::<u64>(),
}
}
pub fn vectorization_factor(&self) -> usize {
match self {
Item::Scalar(_) => 1,
Item::Vector(_, vector_size) => *vector_size,
Item::Atomic(inner)
| Item::Pointer(inner, _)
| Item::Array(inner, _)
| Item::DynamicArray(inner) => inner.vectorization_factor(),
}
}
pub fn with_elem(self, elem: Elem) -> Self {
match self {
Item::Scalar(_) => Item::Scalar(elem),
Item::Vector(_, vector_size) => Item::Vector(elem, vector_size),
Item::Atomic(inner) => Item::Atomic(inner.with_elem(elem).intern()),
Item::Pointer(inner, class) => Item::Pointer(inner.with_elem(elem).intern(), class),
Item::Array(inner, size) => Item::Array(inner.with_elem(elem).intern(), size),
Item::DynamicArray(inner) => Item::DynamicArray(inner.with_elem(elem).intern()),
}
}
pub fn is_ptr(&self) -> bool {
matches!(self, Item::Pointer(..))
}
pub fn fmt_cast_to(&self, item: Item, text: String) -> String {
if *self != item {
format!("{item}({text})")
} else {
text
}
}
}
impl Elem {
pub fn size(&self) -> usize {
match self {
Self::F16 => core::mem::size_of::<half::f16>(),
Self::F32 => core::mem::size_of::<f32>(),
Self::F64 => core::mem::size_of::<f64>(),
Self::I32 => core::mem::size_of::<i32>(),
Self::I64 => core::mem::size_of::<i64>(),
Self::U32 => core::mem::size_of::<u32>(),
Self::U64 => core::mem::size_of::<u64>(),
Self::Bool => core::mem::size_of::<bool>(),
}
}
}
impl Display for Elem {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::F16 => f.write_str("f16"),
Self::F32 => f.write_str("f32"),
Self::F64 => f.write_str("f64"),
Self::I32 => f.write_str("i32"),
Self::I64 => f.write_str("i64"),
Self::U32 => f.write_str("u32"),
Self::U64 => f.write_str("u64"),
Self::Bool => f.write_str("bool"),
}
}
}
impl Display for Item {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Item::Scalar(elem) => write!(f, "{elem}"),
Item::Vector(elem, vector_size) => write!(f, "vec{vector_size}<{elem}>"),
Item::Atomic(inner) => write!(f, "atomic<{inner}>"),
Item::Pointer(inner, class) => match class {
PointerClass::Global(Visibility::Uniform) => write!(f, "ptr<uniform, {inner}>"),
PointerClass::Global(Visibility::Read) => write!(f, "ptr<storage, {inner}, read>"),
PointerClass::Global(Visibility::ReadWrite) => {
write!(f, "ptr<storage, {inner}, read_write>")
}
PointerClass::Shared => write!(f, "ptr<workgroup, {inner}>"),
PointerClass::Local => write!(f, "ptr<function, {inner}>"),
},
Item::Array(inner, size) => {
write!(f, "array<{inner}, {size}>")
}
Item::DynamicArray(inner) => {
write!(f, "array<{inner}>")
}
}
}
}
impl Display for Value {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Value::Value { id, .. } => write!(f, "val_{id}"),
Value::Constant(val, item) => {
match (val, item.elem()) {
(ConstantValue::UInt(v), Elem::U64) if *v > i64::MAX as u64 => {
let as_i64 = *v as i64;
if as_i64 == i64::MIN {
write!(f, "bitcast<u64>(i64(-9223372036854775807) - 1)")
} else {
write!(f, "bitcast<u64>(i64({as_i64}))")
}
}
(ConstantValue::Int(v), Elem::I64) if *v == i64::MIN => {
write!(f, "(i64(-9223372036854775807) - 1)")
}
(_, Elem::U64) | (_, Elem::I64) => write!(f, "{item}({val})"),
_ => write!(f, "{item}({val})"),
}
}
}
}
}
impl Display for Builtin {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Builtin::Id => f.write_str("id"),
Builtin::LocalInvocationIndex => f.write_str("local_idx"),
Builtin::LocalInvocationIdX => f.write_str("local_invocation_id.x"),
Builtin::LocalInvocationIdY => f.write_str("local_invocation_id.y"),
Builtin::LocalInvocationIdZ => f.write_str("local_invocation_id.z"),
Builtin::WorkgroupId => f.write_str("workgroup_id_no_axis"),
Builtin::WorkgroupIdX => f.write_str("workgroup_id.x"),
Builtin::WorkgroupIdY => f.write_str("workgroup_id.y"),
Builtin::WorkgroupIdZ => f.write_str("workgroup_id.z"),
Builtin::GlobalInvocationIdX => f.write_str("global_id.x"),
Builtin::GlobalInvocationIdY => f.write_str("global_id.y"),
Builtin::GlobalInvocationIdZ => f.write_str("global_id.z"),
Builtin::WorkgroupSizeX => f.write_str("WORKGROUP_SIZE_X"),
Builtin::WorkgroupSizeY => f.write_str("WORKGROUP_SIZE_Y"),
Builtin::WorkgroupSizeZ => f.write_str("WORKGROUP_SIZE_Z"),
Builtin::NumWorkgroupsX => f.write_str("num_workgroups.x"),
Builtin::NumWorkgroupsY => f.write_str("num_workgroups.y"),
Builtin::NumWorkgroupsZ => f.write_str("num_workgroups.z"),
Builtin::WorkgroupSize => f.write_str("workgroup_size_no_axis"),
Builtin::NumWorkgroups => f.write_str("num_workgroups_no_axis"),
Builtin::SubgroupSize => f.write_str("subgroup_size"),
Builtin::SubgroupId => f.write_str("subgroup_id"),
Builtin::SubgroupInvocationId => f.write_str("subgroup_invocation_id"),
}
}
}
impl Display for IndexedValue {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let val = &self.val;
let item = val.item();
let index = self.index;
match &self.val {
val if matches!(item, Item::Scalar(_)) => write!(f, "{val}"),
val => write!(f, "{val}[{index}]"),
}
}
}
impl Value {
pub fn fmt_left(&self) -> String {
match self {
Value::Value { .. } => {
format!("let {self}")
}
val => format!("{val}"),
}
}
}
impl IndexedValue {
pub fn fmt_left(&self) -> String {
let item = self.val.item();
match &self.val {
val if matches!(item, Item::Scalar(_)) => val.fmt_left(),
_ => format!("{self}"),
}
}
pub fn fmt_cast(&self, item: Item) -> String {
if self.val.item() != item {
format!("{item}({self})")
} else {
format!("{self}")
}
}
}