use std::num::{NonZeroU32, NonZeroUsize};
use mircuda::DeviceInfo;
use crate::{Error, Result};
mod attention;
mod dense;
mod moe;
mod output;
#[cfg(test)]
mod tests;
pub use dense::{DenseExecution, DensePlan, DensePlanRequest, DenseRole};
pub use moe::{MoeExecution, MoePlan, MoePlanRequest, MoeQuantization};
pub use output::{OutputHeadExecution, OutputHeadPlan, OutputHeadPlanRequest};
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum ExecutionPhase {
Decode,
Prefill,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaNumericalPolicy {
#[default]
Validated,
Throughput,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaKernelAdmission {
#[default]
Stable,
Experimental,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub struct CudaPlanningPolicy {
pub attention: CudaAttentionPolicy,
pub numerical: CudaNumericalPolicy,
pub admission: CudaKernelAdmission,
pub dense_vectors: CudaDenseVectorPolicy,
pub dense_weights: CudaDenseWeightPolicy,
pub moe_fusion: CudaMoeFusionPolicy,
pub moe_batch: CudaMoeBatchPolicy,
pub output_head: CudaOutputHeadPolicy,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaAttentionPolicy {
#[default]
Auto,
Direct,
SplitKv {
partition_tokens: usize,
threshold_tokens: usize,
},
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaDenseWeightPolicy {
#[default]
Bf16,
BlockFp8Role(DenseRole),
Fp8Int4Role(DenseRole),
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaDenseVectorPolicy {
#[default]
Disabled,
Tuned,
Role(DenseRole),
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaMoeFusionPolicy {
#[default]
Disabled,
Tuned,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaMoeBatchPolicy {
#[default]
Auto,
W4A4,
W4A4Direct,
W4A4Hybrid,
W4A4Bucketed,
W4A16,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum CudaOutputHeadPolicy {
#[default]
Auto,
Bf16,
Fp8Blockwise,
Fp8Vectorized,
Fp8Residual,
Fp8BlockVectorized,
Fp8BlockRefined,
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum CudaMemoryArchitecture {
Unified,
Discrete,
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct CudaHardwareProfile {
compute_capability: (u32, u32),
multiprocessor_count: NonZeroU32,
total_memory: NonZeroUsize,
memory_architecture: CudaMemoryArchitecture,
}
impl CudaHardwareProfile {
pub fn new(
compute_capability: (u32, u32),
multiprocessor_count: u32,
total_memory: usize,
memory_architecture: CudaMemoryArchitecture,
) -> Result<Self> {
if compute_capability.0 == 0 {
return Err(Error::InvalidExecutionPlan("CUDA compute capability is missing"));
}
Ok(Self {
compute_capability,
multiprocessor_count: NonZeroU32::new(multiprocessor_count)
.ok_or(Error::InvalidExecutionPlan("CUDA device has no SMs"))?,
total_memory: NonZeroUsize::new(total_memory)
.ok_or(Error::InvalidExecutionPlan("CUDA device has no memory"))?,
memory_architecture,
})
}
pub(super) fn from_device(device: &DeviceInfo) -> Result<Self> {
Self::new(
(
u32::try_from(device.compute_capability.0)?,
u32::try_from(device.compute_capability.1)?,
),
device.multiprocessor_count,
device.total_memory,
if device.integrated {
CudaMemoryArchitecture::Unified
} else {
CudaMemoryArchitecture::Discrete
},
)
}
#[must_use]
pub const fn compute_capability(self) -> (u32, u32) {
self.compute_capability
}
#[must_use]
pub const fn multiprocessor_count(self) -> NonZeroU32 {
self.multiprocessor_count
}
#[must_use]
pub const fn total_memory(self) -> NonZeroUsize {
self.total_memory
}
#[must_use]
pub const fn memory_architecture(self) -> CudaMemoryArchitecture {
self.memory_architecture
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum PlanSource {
Tuned,
ExplicitPolicy,
Fallback,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct CudaExecutionPlanner {
hardware: CudaHardwareProfile,
policy: CudaPlanningPolicy,
}
impl CudaExecutionPlanner {
#[must_use]
pub const fn new(hardware: CudaHardwareProfile, policy: CudaPlanningPolicy) -> Self {
Self { hardware, policy }
}
#[must_use]
pub const fn hardware(self) -> CudaHardwareProfile {
self.hardware
}
#[must_use]
pub const fn policy(self) -> CudaPlanningPolicy {
self.policy
}
}
pub use attention::{AttentionExecution, AttentionPlan, AttentionPlanRequest};