use std::sync::Arc;
use onnx_runtime_ep_api::{
Cost, DeviceBuffer, EpConfig, EpError, ExecutionProvider, Fence, Kernel, KernelMatch,
OpRegistry, Result, deny,
};
use onnx_runtime_ir::{DataType, DeviceId, DeviceType, Node, Shape, TensorLayout};
use crate::kernels::build_cuda_registry_with_metrics;
use crate::kernels::csa_checkpoint::CsaMetrics;
use crate::optimizer::cuda_optimization_passes;
use crate::runtime::{CudaRuntime, cuptr, raw_ptr};
pub struct CudaExecutionProvider {
device: DeviceId,
runtime: Arc<CudaRuntime>,
initialized: bool,
registry: OpRegistry,
csa_metrics: Arc<CsaMetrics>,
}
impl std::fmt::Debug for CudaExecutionProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CudaExecutionProvider")
.field("device", &self.device)
.field("initialized", &self.initialized)
.field("registered_ops", &self.registry.len())
.finish()
}
}
impl CudaExecutionProvider {
pub fn new(ordinal: u32) -> Result<Self> {
let runtime = Arc::new(CudaRuntime::new(ordinal)?);
let csa_metrics = Arc::new(CsaMetrics::default());
let registry = build_cuda_registry_with_metrics(runtime.clone(), csa_metrics.clone());
Ok(Self {
device: DeviceId::cuda(ordinal),
runtime,
initialized: false,
registry,
csa_metrics,
})
}
pub fn initialized(ordinal: u32) -> Result<Self> {
let mut provider = Self::new(ordinal)?;
<Self as ExecutionProvider>::initialize(&mut provider, &EpConfig::default())?;
Ok(provider)
}
pub fn new_default() -> Result<Self> {
Self::new(0)
}
pub fn registry(&self) -> &OpRegistry {
&self.registry
}
pub fn runtime(&self) -> &Arc<CudaRuntime> {
&self.runtime
}
pub fn csa_metrics(&self) -> &Arc<CsaMetrics> {
&self.csa_metrics
}
}
impl ExecutionProvider for CudaExecutionProvider {
fn name(&self) -> &str {
"cuda_ep"
}
fn device_type(&self) -> DeviceType {
DeviceType::Cuda
}
fn device_id(&self) -> DeviceId {
self.device
}
fn initialize(&mut self, _config: &EpConfig) -> Result<()> {
self.runtime.bind()?;
self.initialized = true;
Ok(())
}
fn shutdown(&mut self) -> Result<()> {
self.initialized = false;
Ok(())
}
fn supports_op(
&self,
op: &Node,
opset: u64,
shapes: &[Shape],
input_dtypes: &[DataType],
_layouts: &[TensorLayout],
) -> KernelMatch {
if !self.registry.supports(&op.op_type, &op.domain, opset) {
let domain = if op.domain.is_empty() {
"ai.onnx"
} else {
&op.domain
};
if let Some(since) = self
.registry
.earliest_since_version(&op.op_type, &op.domain)
{
deny!(
"no handler for {}::{} at opset {} — this EP registers {} since opset {} (or: add a claim+handler)",
domain,
op.op_type,
opset,
op.op_type,
since
);
}
deny!(
"no handler for {}::{} at opset {} — add a claim+handler",
domain,
op.op_type,
opset
);
}
if matches!(op.op_type.as_str(), "FusedMatMulBias" | "FusedGemm")
&& op.domain == "com.microsoft"
&& let Some(reason) = crate::kernels::fused_gemm::unsupported_reason(op, shapes)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "BlockQuantizedMatMul"
&& op.domain == "pkg.nxrt"
&& let Some(reason) = crate::kernels::block_quantized_matmul::unsupported_reason(op)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "CompressedSparseAttention"
&& op.domain == "pkg.nxrt"
&& let Some(reason) = crate::kernels::compressed_sparse_attention::unsupported_reason(
op,
shapes,
input_dtypes,
)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "IndexShare"
&& op.domain == "pkg.nxrt"
&& let Some(reason) =
crate::kernels::index_share::unsupported_reason(op, shapes, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "QMoE"
&& op.domain == "com.microsoft"
&& let Some(reason) = crate::kernels::qmoe::unsupported_reason(op)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "Attention"
&& (op.domain.is_empty() || op.domain == "ai.onnx")
&& let Some(reason) =
crate::kernels::standard_attention::unsupported_reason(opset, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "RotaryEmbedding"
&& (op.domain.is_empty() || op.domain == "ai.onnx")
&& let Some(reason) = crate::kernels::rotary_embedding::unsupported_reason(input_dtypes)
{
return KernelMatch::unsupported(reason);
}
if (op.domain.is_empty() || op.domain == "ai.onnx")
&& let Some(reason) =
crate::kernels::standard_claims::unsupported_reason(op, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
if matches!(
op.op_type.as_str(),
"Equal" | "Greater" | "Less" | "GreaterOrEqual" | "LessOrEqual"
) && (op.domain.is_empty() || op.domain == "ai.onnx")
&& let Some(reason) =
crate::kernels::pointwise::comparison_unsupported_reason(&op.op_type, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
let output_layouts = vec![TensorLayout::contiguous(); op.outputs.len()];
let elems: u64 = shapes
.iter()
.map(|s| {
s.iter()
.map(|d| d.as_static().unwrap_or(1) as u64)
.product::<u64>()
})
.sum();
let cost = Cost::new(elems as f64 * 0.01, elems as f64 * 0.01, 0.0)
.with_launch_us(10.0)
.with_bytes_moved(elems.saturating_mul(4));
KernelMatch::Supported {
cost,
required_input_layouts: None,
output_layouts,
}
}
fn get_kernel(&self, op: &Node, shapes: &[Vec<usize>], opset: u64) -> Result<Box<dyn Kernel>> {
let factory = self
.registry
.lookup(&op.op_type, &op.domain, opset)
.ok_or_else(|| EpError::NoEpForOp {
domain: if op.domain.is_empty() {
"ai.onnx".to_string()
} else {
op.domain.clone()
},
op_type: op.op_type.clone(),
opset,
})?;
factory.create(op, shapes)
}
fn custom_passes(&self) -> Vec<Box<dyn onnx_runtime_optimizer::OptimizationPass>> {
cuda_optimization_passes()
}
fn allocate(&self, size: usize, alignment: usize) -> Result<DeviceBuffer> {
if alignment == 0 || !alignment.is_power_of_two() {
return Err(EpError::AlignmentError);
}
let dptr = self.runtime.alloc_raw(size)?;
Ok(unsafe { DeviceBuffer::from_raw_parts(raw_ptr(dptr), self.device, size, alignment) })
}
fn deallocate(&self, buffer: DeviceBuffer) -> Result<()> {
assert_eq!(
buffer.device(),
self.device,
"cuda_ep: refusing to deallocate a buffer from device {:?}",
buffer.device()
);
if buffer.is_borrowed() {
return Ok(());
}
let dptr = cuptr(buffer.into_raw());
unsafe { self.runtime.free_raw(dptr) }
}
fn copy(&self, src: &DeviceBuffer, dst: &mut DeviceBuffer, size: usize) -> Result<()> {
assert_eq!(
src.device(),
self.device,
"cuda_ep::copy: foreign src buffer"
);
assert_eq!(
dst.device(),
self.device,
"cuda_ep::copy: foreign dst buffer"
);
if size > src.len() || size > dst.len() {
return Err(EpError::KernelFailed(format!(
"cuda_ep::copy: size {size} exceeds src {} or dst {}",
src.len(),
dst.len()
)));
}
if size == 0 {
return Ok(());
}
let src_p = cuptr(src.as_ptr());
let dst_p = cuptr(dst.as_mut_ptr());
unsafe { self.runtime.dtod(src_p, dst_p, size) }
}
fn copy_async(&self, src: &DeviceBuffer, dst: &mut DeviceBuffer, size: usize) -> Result<Fence> {
self.copy(src, dst, size)?;
Ok(Fence::default())
}
fn device_argmax_supported(&self) -> bool {
true
}
fn device_argmax(
&self,
logits: &DeviceBuffer,
elements: usize,
dtype: DataType,
result: &mut DeviceBuffer,
) -> Result<()> {
crate::kernels::device_argmax::launch(&self.runtime, logits, elements, dtype, result)
}
fn copy_from_host(&self, src: &[u8], dst: &mut DeviceBuffer) -> Result<()> {
assert_eq!(
dst.device(),
self.device,
"cuda_ep::copy_from_host: foreign dst buffer"
);
if src.len() > dst.len() {
return Err(EpError::KernelFailed(format!(
"cuda_ep::copy_from_host: source {} bytes exceeds dst {}",
src.len(),
dst.len()
)));
}
if src.is_empty() {
return Ok(());
}
unsafe { self.runtime.htod(src, cuptr(dst.as_mut_ptr())) }
}
fn copy_from_host_at(
&self,
src: &[u8],
dst: &mut DeviceBuffer,
byte_offset: usize,
) -> Result<()> {
assert_eq!(
dst.device(),
self.device,
"cuda_ep::copy_from_host_at: foreign dst buffer"
);
let end = byte_offset.checked_add(src.len()).ok_or_else(|| {
EpError::KernelFailed("cuda_ep::copy_from_host_at: upload range overflows".into())
})?;
if end > dst.len() {
return Err(EpError::KernelFailed(format!(
"cuda_ep::copy_from_host_at: range {byte_offset}..{end} exceeds dst {}",
dst.len()
)));
}
if src.is_empty() {
return Ok(());
}
let ptr = cuptr(dst.as_mut_ptr())
.checked_add(byte_offset as u64)
.ok_or_else(|| {
EpError::KernelFailed(
"cuda_ep::copy_from_host_at: device pointer offset overflows".into(),
)
})?;
unsafe { self.runtime.htod(src, ptr) }
}
fn copy_to_host(&self, src: &DeviceBuffer, dst: &mut [u8]) -> Result<()> {
assert_eq!(
src.device(),
self.device,
"cuda_ep::copy_to_host: foreign src buffer"
);
if dst.len() > src.len() {
return Err(EpError::KernelFailed(format!(
"cuda_ep::copy_to_host: destination {} bytes exceeds src {}",
dst.len(),
src.len()
)));
}
if dst.is_empty() {
return Ok(());
}
unsafe { self.runtime.dtoh(dst, cuptr(src.as_ptr())) }
}
fn begin_device_graph_capture(&self, kernels: &[&dyn Kernel]) -> Result<()> {
self.runtime.begin_graph_capture(kernels)
}
fn end_device_graph_capture(&self) -> Result<()> {
self.runtime.end_graph_capture()
}
fn abort_device_graph_capture(&self) -> Result<()> {
self.runtime.abort_graph_capture()
}
fn replay_device_graph(&self) -> Result<()> {
self.runtime.replay_graph()
}
fn replay_device_graph_segment(&self, index: usize) -> Result<()> {
self.runtime.replay_graph_segment(index)
}
fn reset_device_graph(&self) -> Result<bool> {
let invalidated = self.runtime.reset_graph()?;
self.runtime.reset_capture_error()?;
Ok(invalidated)
}
fn check_device_capture_error(&self) -> Result<u32> {
self.runtime.check_capture_error()
}
fn device_allocation_counts(&self) -> Option<(u64, u64)> {
let counts = self.runtime.allocation_counts();
Some((counts.allocations, counts.frees))
}
fn sync(&self) -> Result<()> {
self.runtime.synchronize()
}
}