#![allow(unused)]
use core::{fmt, ops::Deref};
use cubecl_core::{
self as cubecl,
ir::types::Fp8Format,
ir::{
dialect::general::CastOp,
interfaces::{ScalarType, TypedExt},
match_ty,
prelude::*,
types::{VectorType, scalar::*},
},
prelude::*,
};
use pliron::{printable::Printable, utils::apfloat::Float8E5M2};
use crate::{
cuda::{cuda_op_with_out, ty::*},
shared::{
CppValue,
lowering::LowerOp,
ty::{TypeExt, TypeExtCPP, TypedExtCPP},
},
target::Cuda,
};
#[op_interface_impl]
impl LowerOp<Cuda> for CastOp {
fn should_lower(&self, ctx: &Context) -> bool {
let input = self.input(ctx);
let out = self.get_result(ctx);
let should_lower_from = (input.is_fp8_fp6_fp4(ctx) || input.is_float4x2(ctx))
&& intermediate_for_ty(ctx, input.get_type(ctx)) != out.get_type(ctx);
let should_lower_to = (out.is_fp8_fp6_fp4(ctx) || out.is_float4x2(ctx))
&& !encodes_directly(ctx, input, out.get_type(ctx));
should_lower_from || should_lower_to
}
fn lower(&self, scope: &Scope) -> Vec<Value> {
let ctx = scope.ctx();
let mut current = self.input(ctx);
let out_ty = self.get_result(ctx).get_type(ctx);
if current.is_fp8_fp6_fp4(ctx) || current.is_float4x2(ctx) {
let intermediate = intermediate_for_ty(ctx, current.get_type(ctx));
current = cast_value(scope, current, intermediate);
}
if (out_ty.is_fp8_fp6_fp4(ctx) || out_ty.is_float4x2(ctx))
&& !encodes_directly(ctx, current, out_ty)
{
let intermediate = match is_fp8(ctx, out_ty) {
true => f32_like(ctx, out_ty),
false => intermediate_for_ty(ctx, out_ty),
};
current = cast_value(scope, current, intermediate);
}
vec![cast_value(scope, current, out_ty)]
}
}
fn encodes_directly(ctx: &Context, input: Value, out_ty: TypeHandle) -> bool {
if !is_fp8(ctx, out_ty) {
return intermediate_for_ty(ctx, out_ty) == input.get_type(ctx);
}
let scalar = input.get_type(ctx).scalar_ty(ctx);
scalar.is_float16(ctx)
|| scalar.is_bfloat16(ctx)
|| scalar.is_float32(ctx)
|| scalar.is_float64(ctx)
}
fn is_fp8(ctx: &Context, ty: TypeHandle) -> bool {
Fp8Format::of_type(ctx, ty.scalar_ty(ctx)).is_some()
}
fn f32_like(ctx: &Context, ty: TypeHandle) -> TypeHandle {
vectorized(ctx, Float32Type::get(ctx).to_handle(), ty.vector_size(ctx))
}
fn intermediate_for_ty(ctx: &Context, ty: TypeHandle) -> TypeHandle {
let vector_size = ty.vector_size(ctx);
let intermediate = if ty.scalar_ty(ctx).deref(ctx).is::<Float8E8M0Type>() {
BFloat16Type::get(ctx).to_handle()
} else if ty.is_float4x2(ctx) {
return VectorType::get(ctx, Float16Type::get(ctx).to_handle(), vector_size * 2)
.to_handle();
} else {
Float16Type::get(ctx).to_handle()
};
vectorized(ctx, intermediate, vector_size)
}
fn vectorized(ctx: &Context, scalar: TypeHandle, vector_size: usize) -> TypeHandle {
if vector_size > 1 {
VectorType::get(ctx, scalar, vector_size).to_handle()
} else {
scalar
}
}
cuda_op_with_out!(CastOp, |op, ctx| {
let input = op.input(ctx);
let input_name = input.name(ctx);
let out_ty = op.get_result(ctx).get_type(ctx);
let out_scalar = out_ty.scalar_ty(ctx).deref(ctx);
if out_scalar.is::<Complex32Type>() {
if input.is_complex(ctx) {
format!("make_cuFloatComplex({input_name}.x, {input_name}.y)")
} else {
format!("make_cuFloatComplex({input_name}, 0.0f)")
}
} else if out_scalar.is::<Complex64Type>() {
if input.is_complex(ctx) {
format!("make_cuDoubleComplex({input_name}.x, {input_name}.y)")
} else {
format!("make_cuDoubleComplex({input_name}, 0.0)")
}
} else if input.is_complex(ctx) {
if out_ty.is_tfloat32(ctx) {
format!("nvcuda::wmma::__float_to_tf32({input_name}.x)")
} else if out_ty.is_bool(ctx) {
format!("({input_name}.x != 0 || {input_name}.y != 0)")
} else {
format!("{}({input_name}.x)", out_ty.to_cpp(ctx))
}
} else if input.is_fp8_fp6_fp4(ctx) || input.is_packed_fp6_fp8_fp4(ctx) {
cast_minifloat_to_half(ctx, input)
} else if out_ty.is_fp8_fp6_fp4(ctx) || out_ty.is_packed_fp6_fp8_fp4(ctx) {
cast_half_to_minifloat(ctx, input, out_ty)
} else if out_ty.is_tfloat32(ctx) {
format!("nvcuda::wmma::__float_to_tf32({input_name})")
} else {
format!("{}({input_name})", out_ty.to_cpp(ctx))
}
});
fn cast_minifloat_to_half(ctx: &Context, input: Value) -> String {
let in_ty = input.get_type(ctx).deref(ctx);
let in_val = input.name(ctx);
match_ty!((in_ty) {
Float8E8M0Type => format!("__nv_bfloat16(__nv_cvt_e8m0_to_bf16raw({in_val}))"),
Float8E8M0x2Type => format!("__nv_bfloat162(__nv_cvt_e8m0x2_to_bf162raw({in_val}))"),
Float8E4M3Type => format!("__half(__nv_cvt_fp8_to_halfraw({in_val}, __NV_E4M3))"),
Float8E4M3x2Type => format!("__half2(__nv_cvt_fp8x2_to_halfraw2({in_val}, __NV_E4M3))"),
Float8E5M2Type => format!("__half(__nv_cvt_fp8_to_halfraw({in_val}, __NV_E5M2))"),
Float8E5M2x2Type => format!("__half2(__nv_cvt_fp8x2_to_halfraw2({in_val}, __NV_E5M2))"),
Float6E2M3Type => format!("__half(__nv_cvt_fp6_to_halfraw({in_val}, __NV_E2M3))"),
Float6E2M3x2Type => format!("__half2(__nv_cvt_fp6x2_to_halfraw2({in_val}, __NV_E2M3))"),
Float6E3M2Type => format!("__half(__nv_cvt_fp6_to_halfraw({in_val}, __NV_E3M2))"),
Float6E3M2x2Type => format!("__half(__nv_cvt_fp6x2_to_halfraw2({in_val}, __NV_E3M2))"),
Float4E2M1Type => format!("__half(__nv_cvt_fp4_to_halfraw({in_val}, __NV_E2M1))"),
Float4E2M1x2Type => format!("__half2(__nv_cvt_fp4x2_to_halfraw2({in_val}, __NV_E2M1))"),;
_ => panic!("Unsupported type {}", in_ty.display(ctx))
})
}
fn cast_half_to_minifloat(ctx: &Context, input: Value, out_ty: TypeHandle) -> String {
let in_val = input.name(ctx);
let fp8_source = || fp8_source_prefix(ctx, input);
match_ty!((out_ty.deref(ctx)) {
Float8E8M0Type => format!("__nv_cvt_bfloat16raw_to_e8m0({in_val}, __NV_NOSAT, cudaRoundPosInf)"),
Float8E8M0x2Type => format!("__nv_cvt_bfloat162raw_to_e8m0x2({in_val}, __NV_NOSAT, cudaRoundPosInf)"),
Float8E4M3Type => format!("__nv_cvt_{}_to_fp8({in_val}, __NV_SATFINITE, __NV_E4M3)", fp8_source()),
Float8E4M3x2Type => format!("__nv_cvt_{}_to_fp8x2({in_val}, __NV_SATFINITE, __NV_E4M3)", fp8_source()),
Float8E5M2Type => format!("__nv_cvt_{}_to_fp8({in_val}, __NV_SATFINITE, __NV_E5M2)", fp8_source()),
Float8E5M2x2Type => format!("__nv_cvt_{}_to_fp8x2({in_val}, __NV_SATFINITE, __NV_E5M2)", fp8_source()),
Float6E2M3Type => format!("__nv_cvt_halfraw_to_fp6({in_val}, __NV_E2M3, cudaRoundNearest)"),
Float6E2M3x2Type => format!("__nv_cvt_halfraw2_to_fp6x2({in_val}, __NV_E2M3, cudaRoundNearest)"),
Float6E3M2Type => format!("__nv_cvt_halfraw_to_fp6({in_val}, __NV_E3M2, cudaRoundNearest)"),
Float6E3M2x2Type => format!("__nv_cvt_halfraw2_to_fp6x2({in_val}, __NV_E3M2, cudaRoundNearest)"),
Float4E2M1Type => format!("__nv_cvt_halfraw_to_fp4({in_val}, __NV_E2M1, cudaRoundNearest)"),
Float4E2M1x2Type => format!("__nv_cvt_halfraw2_to_fp4x2({in_val}, __NV_E2M1, cudaRoundNearest)"),;
_ => panic!("Unsupported type {}", out_ty.deref(ctx).display(ctx))
})
}
fn fp8_source_prefix(ctx: &Context, input: Value) -> &'static str {
let ty = input.get_type(ctx);
let unsupported = || {
panic!(
"fp8 converts from a float scalar or a packed 16-bit pair, got {}",
ty.deref(ctx).display(ctx)
)
};
match_ty!((ty.deref(ctx)) {
Float16Type => "halfraw",
Float16x2Type => "halfraw2",
BFloat16Type => "bfloat16raw",
BFloat16x2Type => "bfloat16raw2",
Float32Type => "float",
Float64Type => "double",;
_ => unsupported()
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
shared::operation::OpToCPP,
target::{CtxTarget, Target},
};
use cubecl_core::ir::{ComplexKind, ConstantValue, ElemType, FloatKind, UIntKind};
use pliron::{attribute::boxed_attr_cast, builtin::ops::ConstantOp};
fn cast(input: ConstantValue, input_ty: ElemType, output_ty: ElemType) -> String {
let mut ctx = Context::new();
ctx.set_target(Target::Cuda);
let input_attr = input.as_attribute(&ctx, input_ty);
let input_attr = boxed_attr_cast(input_attr).unwrap();
let input = ConstantOp::new(&mut ctx, input_attr).get_result(&ctx);
let input_name = input.name(&ctx).to_string();
let output_ty = output_ty.to_type(&ctx);
let op = CastOp::new(&mut ctx, output_ty, input);
let cpp = OpToCPP::<Cuda>::to_cpp(&op, &ctx);
cpp.split_once(" = ")
.unwrap()
.1
.replace(&input_name, "input")
}
#[test]
fn complex_casts_use_cucomplex_components_and_constructors() {
assert_eq!(
cast(
ConstantValue::UInt(0),
UIntKind::U32.into(),
ComplexKind::C64.into()
),
"make_cuDoubleComplex(input, 0.0);\n"
);
assert_eq!(
cast(
ConstantValue::Float(1.0),
FloatKind::F64.into(),
ComplexKind::C32.into()
),
"make_cuFloatComplex(input, 0.0f);\n"
);
assert_eq!(
cast(
ConstantValue::Complex(1.0, 2.0),
ComplexKind::C64.into(),
ComplexKind::C32.into()
),
"make_cuFloatComplex(input.x, input.y);\n"
);
assert_eq!(
cast(
ConstantValue::Complex(1.0, 2.0),
ComplexKind::C32.into(),
ComplexKind::C64.into()
),
"make_cuDoubleComplex(input.x, input.y);\n"
);
assert_eq!(
cast(
ConstantValue::Complex(1.0, 2.0),
ComplexKind::C32.into(),
FloatKind::F64.into()
),
"double(input.x);\n"
);
assert_eq!(
cast(
ConstantValue::Complex(0.0, 1.0),
ComplexKind::C32.into(),
ElemType::Bool
),
"(input.x != 0 || input.y != 0);\n"
);
}
}