libmir-cuda 0.3.0

CUDA inference backend for libmir
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
pub use architecture::cuda::CudaDecoderPlan;

use crate::{Error, Result, kernels::QkvNormalization as KernelNormalization};

#[cfg(test)]
mod tests;

pub fn graph_normalization(plan: &CudaDecoderPlan) -> Result<KernelNormalization> {
    let normalization = match architecture::cuda::graph_normalization(plan) {
        Ok(normalization) => normalization,
        Err(error) => return Err(Error::UnsupportedDecoderLayer(error.to_string())),
    };
    Ok(match normalization {
        architecture::cuda::CudaQkvNormalization::None => KernelNormalization::NONE,
        architecture::cuda::CudaQkvNormalization::All => KernelNormalization::ALL,
        architecture::cuda::CudaQkvNormalization::QueryKey => KernelNormalization::QUERY_KEY,
    })
}