Skip to main content

ferrum_kernels/
lib.rs

1//! Ferrum unified compute kernels for high-performance inference.
2//!
3//! Provides the `Backend` trait and implementations for CUDA, Metal, and CPU.
4//! On CUDA builds, kernels are compiled to PTX during `cargo build` and loaded
5//! on demand at runtime.
6
7pub fn configure_native_profile_sink(
8    config: &ferrum_bench_core::ProfileSinkConfig,
9) -> std::io::Result<()> {
10    #[cfg(all(feature = "cuda", feature = "vllm-moe-marlin"))]
11    backend::cuda::marlin::configure_vllm_moe_profile_sink(config)?;
12    #[cfg(not(all(feature = "cuda", feature = "vllm-moe-marlin")))]
13    let _ = config;
14    Ok(())
15}
16
17#[cfg(feature = "cuda")]
18pub fn cuda_device_count() -> Result<usize, String> {
19    cudarc::driver::CudaContext::device_count()
20        .map_err(|error| format!("failed to query CUDA device count: {error}"))
21        .and_then(|count| {
22            usize::try_from(count)
23                .map_err(|_| format!("CUDA driver returned a negative device count: {count}"))
24        })
25}
26
27#[cfg(feature = "cuda")]
28pub fn cuda_device_name(ordinal: usize) -> Result<String, String> {
29    cudarc::driver::CudaContext::new(ordinal)
30        .map_err(|error| format!("failed to open CUDA device {ordinal}: {error}"))?
31        .name()
32        .map_err(|error| format!("failed to query CUDA device {ordinal} name: {error}"))
33}
34
35pub mod backend;
36pub use backend::probe_device_memory;
37pub mod gguf_blocks;
38#[cfg(test)]
39pub(crate) mod hadamard;
40pub mod native_ops;
41
42pub mod linear;
43pub use linear::{Linear, LinearMetadata, LinearProjectionRole};
44
45pub mod stacked_expert;
46pub use stacked_expert::StackedExpertGgufLinear;
47
48pub mod marlin_expert_stack;
49pub mod marlin_fp8_materializer;
50pub mod marlin_repack;
51pub mod mxfp4_marlin_materializer;
52pub use marlin_expert_stack::MarlinExpertStack;
53
54pub mod quant_linear;
55
56pub mod attention;
57
58pub mod moe_host;
59
60// Audit #9: Metal GGUF k-quant kernels (q4_k_*, q6_k_*, moe_*) physically
61// live in `backend/metal/` now. Re-exported here so external callers'
62// `ferrum_kernels::q4_k_gemm::*` paths + internal `crate::q4_k_*::*` paths
63// keep working unchanged. (`moe_host` stays top-level — it's the CPU
64// reference impl used from `ferrum-models`, not a Metal kernel.)
65#[cfg(all(target_os = "macos", feature = "metal"))]
66pub use backend::metal::{
67    moe_post_ops, moe_post_ops_batched, moe_router, q4_k, q4_k_gemm, q4_k_gemv, q4_k_gemv_v2,
68    q4_k_moe_id_gate_up_silu, q4_k_moe_id_gate_up_silu_batched, q4_k_moe_id_gemm, q4_k_moe_id_gemv,
69    q4_k_moe_id_gemv_batched, q6_k_gemm, q6_k_gemv, q6_k_moe_id_gemm, q6_k_moe_id_gemv,
70    q6_k_moe_id_gemv_batched,
71};
72
73#[cfg(feature = "cuda")]
74pub(crate) mod ptx {
75    // Generated by build.rs from all .cu sources. Some kernels (e.g.
76    // SOFTMAX, BATCHED_FLASH_DECODE_ATTENTION) are emitted unconditionally
77    // but only loaded behind specific code paths, so dead_code fires in
78    // configs that don't hit them.
79    #![allow(dead_code)]
80    include!(concat!(env!("OUT_DIR"), "/ptx.rs"));
81}
82
83// Audit #9: CUDA kernels physically live under `backend/cuda/` now.
84// Re-exports preserve the historical `ferrum_kernels::foo::*` public
85// surface + internal `crate::foo::*` paths.
86//
87// Two files stay at the crate root for now because they would otherwise
88// collide with same-named files already under `backend/cuda/` (which
89// host the Backend-trait impls, not the kernel launchers):
90//   - `int8_kv.rs` (top-level launchers `launch_int8_paged_decode_*`)
91//   - `quant.rs`   (top-level `dequant_int4` legacy path)
92
93#[cfg(feature = "cuda")]
94pub mod int8_kv;
95#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
96pub mod quant;
97
98#[cfg(feature = "cuda")]
99pub use backend::cuda::{cublas, decode_buffers, gpu_paged_kv, marlin};
100
101#[cfg(all(feature = "cuda", not(target_os = "windows")))]
102pub use backend::cuda::nccl_comm;
103
104#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
105pub use backend::cuda::{cuda_decode, cuda_graph, weight_store};
106
107#[cfg(all(
108    feature = "cuda",
109    feature = "candle-cuda-compat",
110    not(target_os = "windows")
111))]
112pub use backend::cuda::tp_decode;
113
114#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
115pub use backend::cuda::decode_attention::decode_attention;
116#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
117pub use backend::cuda::fused_add_rms_norm::fused_add_rms_norm;
118#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
119pub use backend::cuda::fused_silu_mul::fused_silu_mul;
120#[cfg(feature = "cuda")]
121pub use backend::cuda::gated_delta_rule::recurrent_gated_delta_rule_f32;
122#[cfg(feature = "cuda")]
123pub use backend::cuda::linear_attention::{gated_rms_norm_f32, linear_attention_prepare_f32};
124#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
125pub use backend::cuda::residual_add::residual_add;
126#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
127pub use backend::cuda::rms_norm::rms_norm;
128#[cfg(all(feature = "cuda", feature = "candle-cuda-compat"))]
129pub use backend::cuda::rope::rope;
130
131// Preserve `crate::triton_ptx` / `crate::triton_meta` paths for in-crate
132// callers (e.g. `quant_linear::cuda_marlin::CudaMarlinLinear::forward`).
133// These are NOT part of the kernels-crate public API.
134#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
135pub(crate) use backend::cuda::{triton_meta, triton_ptx};
136
137#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
138pub use backend::cuda::triton_add_bias::add_bias_triton;
139#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
140pub use backend::cuda::triton_fused_add_rms_norm::fused_add_rms_norm_triton;
141#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
142pub use backend::cuda::triton_fused_silu_mul::fused_silu_mul_triton;
143#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
144pub use backend::cuda::triton_gelu::gelu_triton;
145#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
146pub use backend::cuda::triton_layer_norm::layer_norm_triton;
147#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
148pub use backend::cuda::triton_residual_add::residual_add_triton;
149#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
150pub use backend::cuda::triton_residual_add_inplace::residual_add_inplace_triton;
151#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
152pub use backend::cuda::triton_rms_norm::rms_norm_triton;
153#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
154pub use backend::cuda::triton_softmax::softmax_triton;
155#[cfg(all(feature = "cuda", feature = "triton-kernels"))]
156pub use backend::cuda::{triton_fused_moe, triton_w4a16};
157
158// vLLM gptq_marlin port (Phase 12). Behind its own feature for opt-in
159// while we validate correctness + perf vs ferrum's existing IST-DASLab Marlin.
160#[cfg(feature = "vllm-marlin")]
161pub use backend::cuda::vllm_marlin;