use mircuda::{
CompileOptions, Compiler, DeviceBuffer, LaunchConfig, Stream, TypedKernel, bf16, cuda_export,
cuda_kernel_file,
};
use super::{DirectFp8Activation, DirectFp8Format, DirectFp8Spec};
use crate::{
Result,
kernels::geometry::{narrow, require},
};
cuda_export!(F32ScaleKernel = "libmir_cuda_direct_fp8_bf16_linear_f32_scale"(
input: &DeviceBuffer<bf16>, weight: &DeviceBuffer<u8>, scales: &DeviceBuffer<f32>,
input_scale: &DeviceBuffer<f32>, bias: &DeviceBuffer<bf16>,
output: &mut DeviceBuffer<bf16>, tokens: u32, rows: u32, columns: u32,
scale_rows: u32, scale_columns: u32, scale_row_size: u32, scale_group_size: u32,
inverse_scale: u32,
has_bias: u32, activation_mode: u32, e5m2: u32,
));
cuda_export!(Bf16ScaleKernel = "libmir_cuda_direct_fp8_bf16_linear_bf16_scale"(
input: &DeviceBuffer<bf16>, weight: &DeviceBuffer<u8>, scales: &DeviceBuffer<bf16>,
input_scale: &DeviceBuffer<bf16>, bias: &DeviceBuffer<bf16>,
output: &mut DeviceBuffer<bf16>, tokens: u32, rows: u32, columns: u32,
scale_rows: u32, scale_columns: u32, scale_row_size: u32, scale_group_size: u32,
inverse_scale: u32,
has_bias: u32, activation_mode: u32, e5m2: u32,
));
#[derive(Clone, Debug)]
pub struct DirectFp8Linear {
f32_scale: TypedKernel<F32ScaleKernel>,
bf16_scale: TypedKernel<Bf16ScaleKernel>,
spec: DirectFp8Spec,
}
#[derive(Clone, Copy, Debug)]
pub struct DirectFp8Scales<'a, T: mircuda::DeviceElement> {
pub weight: &'a DeviceBuffer<T>,
pub activation: &'a DeviceBuffer<T>,
}
impl DirectFp8Linear {
pub fn compile(compiler: &Compiler, spec: DirectFp8Spec) -> Result<Self> {
let module = compiler.compile(
cuda_kernel_file!("../../../kernels/direct_fp8.cu"),
&CompileOptions {
fast_math: false,
..CompileOptions::default()
},
)?;
Ok(Self {
f32_scale: module.kernel()?,
bf16_scale: module.kernel()?,
spec,
})
}
pub fn execute(
&self,
stream: &Stream,
input: &DeviceBuffer<bf16>,
weight: &DeviceBuffer<u8>,
scales: DirectFp8Scales<'_, f32>,
bias: Option<&DeviceBuffer<bf16>>,
output: &mut DeviceBuffer<bf16>,
) -> Result<()> {
self.validate(input, weight, scales.weight.len(), scales.activation, bias, output)?;
let (scale_rows, scale_columns, scale_row_size, scale_group_size) =
self.spec.scale_geometry()?;
Ok(self.f32_scale.launch(
stream,
self.launch()?,
(
input,
weight,
scales.weight,
scales.activation,
bias.unwrap_or(input),
output,
narrow(self.spec.tokens)?,
narrow(self.spec.output_features)?,
narrow(self.spec.input_features)?,
narrow(scale_rows)?,
narrow(scale_columns)?,
narrow(scale_row_size)?,
narrow(scale_group_size)?,
u32::from(self.spec.inverse_scale),
u32::from(bias.is_some()),
self.activation_mode(),
self.e5m2(),
),
)?)
}
pub fn execute_bf16_scales(
&self,
stream: &Stream,
input: &DeviceBuffer<bf16>,
weight: &DeviceBuffer<u8>,
scales: DirectFp8Scales<'_, bf16>,
bias: Option<&DeviceBuffer<bf16>>,
output: &mut DeviceBuffer<bf16>,
) -> Result<()> {
self.validate(input, weight, scales.weight.len(), scales.activation, bias, output)?;
let (scale_rows, scale_columns, scale_row_size, scale_group_size) =
self.spec.scale_geometry()?;
Ok(self.bf16_scale.launch(
stream,
self.launch()?,
(
input,
weight,
scales.weight,
scales.activation,
bias.unwrap_or(input),
output,
narrow(self.spec.tokens)?,
narrow(self.spec.output_features)?,
narrow(self.spec.input_features)?,
narrow(scale_rows)?,
narrow(scale_columns)?,
narrow(scale_row_size)?,
narrow(scale_group_size)?,
u32::from(self.spec.inverse_scale),
u32::from(bias.is_some()),
self.activation_mode(),
self.e5m2(),
),
)?)
}
fn validate<S: mircuda::DeviceElement, T: mircuda::DeviceElement>(
&self,
input: &DeviceBuffer<bf16>,
weight: &DeviceBuffer<u8>,
scales: usize,
input_scale: &DeviceBuffer<S>,
bias: Option<&DeviceBuffer<bf16>>,
output: &DeviceBuffer<T>,
) -> Result<()> {
require("direct FP8 input", self.spec.input_elements()?, input.len())?;
require("direct FP8 weight", self.spec.weight_elements()?, weight.len())?;
require("direct FP8 scales", self.spec.scale_elements()?, scales)?;
require("direct FP8 activation scale", 1, input_scale.len())?;
if let Some(bias) = bias {
require("direct FP8 bias", self.spec.output_features, bias.len())?;
}
require("direct FP8 output", self.spec.output_elements()?, output.len())
}
fn launch(&self) -> Result<LaunchConfig> {
Ok(LaunchConfig {
grid: (narrow(self.spec.output_features.div_ceil(64))?, narrow(self.spec.tokens)?, 1),
block: (256, 1, 1),
shared_memory_bytes: 0,
})
}
fn activation_mode(&self) -> u32 {
match self.spec.activation {
DirectFp8Activation::Bf16 => 0,
DirectFp8Activation::DynamicE4M3Token => 1,
DirectFp8Activation::StaticE4M3Tensor => 2,
}
}
fn e5m2(&self) -> u32 {
u32::from(self.spec.format == DirectFp8Format::E5M2)
}
}