libmir-cuda 0.3.0

CUDA inference backend for libmir
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)]
/// Compiled direct-checkpoint E4M3 or E5M2 projection.
pub struct DirectFp8Linear {
    f32_scale: TypedKernel<F32ScaleKernel>,
    bf16_scale: TypedKernel<Bf16ScaleKernel>,
    spec: DirectFp8Spec,
}

#[derive(Clone, Copy, Debug)]
/// Device-resident weight and activation scales for one direct FP8 launch.
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)
    }
}