use onnx_runtime_ep_api::{DeviceBuffer, ExecutionProvider, Fence, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrefetchStep {
Prefetch { expert: usize, buffer: usize },
Await { expert: usize },
Compute { expert: usize, buffer: usize },
}
pub fn plan_double_buffer(num_experts: usize) -> Vec<PrefetchStep> {
let mut steps = Vec::new();
if num_experts == 0 {
return steps;
}
steps.push(PrefetchStep::Prefetch {
expert: 0,
buffer: 0,
});
for n in 0..num_experts {
if n + 1 < num_experts {
steps.push(PrefetchStep::Prefetch {
expert: n + 1,
buffer: (n + 1) % 2,
});
}
steps.push(PrefetchStep::Await { expert: n });
steps.push(PrefetchStep::Compute {
expert: n,
buffer: n % 2,
});
}
steps
}
pub fn drive_double_buffer<F>(
ep: &dyn ExecutionProvider,
buffers: &mut [DeviceBuffer; 2],
sources: &[DeviceBuffer],
sizes: &[usize],
mut compute: F,
) -> Result<()>
where
F: FnMut(usize, &DeviceBuffer) -> Result<()>,
{
assert_eq!(
sources.len(),
sizes.len(),
"double-buffer prefetch: sources and sizes must have equal length"
);
let num_experts = sources.len();
if num_experts == 0 {
return Ok(());
}
let mut last_compute_fence: [Fence; 2] = [Fence::signalled(), Fence::signalled()];
let mut current_fence: Fence = ep.copy_async(&sources[0], &mut buffers[0], sizes[0])?;
for n in 0..num_experts {
let current_slot = n % 2;
let next_fence = if n + 1 < num_experts {
let next_slot = (n + 1) % 2;
ep.copy_wait_fence(&last_compute_fence[next_slot])?;
Some(ep.copy_async(&sources[n + 1], &mut buffers[next_slot], sizes[n + 1])?)
} else {
None
};
ep.wait_fence(¤t_fence)?;
compute(n, &buffers[current_slot])?;
last_compute_fence[current_slot] = ep.record_compute_fence()?;
if let Some(next) = next_fence {
current_fence = next;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
use onnx_runtime_ep_api::{EpConfig, KernelMatch};
use onnx_runtime_ir::{DataType, DeviceId, DeviceType, Node, Shape, TensorLayout};
#[test]
fn plan_double_buffer_empty_is_empty() {
assert!(plan_double_buffer(0).is_empty());
}
#[test]
fn plan_double_buffer_single_expert_needs_no_second_prefetch() {
assert_eq!(
plan_double_buffer(1),
vec![
PrefetchStep::Prefetch {
expert: 0,
buffer: 0
},
PrefetchStep::Await { expert: 0 },
PrefetchStep::Compute {
expert: 0,
buffer: 0
},
]
);
}
#[test]
fn plan_double_buffer_interleaves_next_prefetch_before_current_compute() {
let steps = plan_double_buffer(4);
assert_eq!(
steps,
vec![
PrefetchStep::Prefetch {
expert: 0,
buffer: 0
},
PrefetchStep::Prefetch {
expert: 1,
buffer: 1
},
PrefetchStep::Await { expert: 0 },
PrefetchStep::Compute {
expert: 0,
buffer: 0
},
PrefetchStep::Prefetch {
expert: 2,
buffer: 0
},
PrefetchStep::Await { expert: 1 },
PrefetchStep::Compute {
expert: 1,
buffer: 1
},
PrefetchStep::Prefetch {
expert: 3,
buffer: 1
},
PrefetchStep::Await { expert: 2 },
PrefetchStep::Compute {
expert: 2,
buffer: 0
},
PrefetchStep::Await { expert: 3 },
PrefetchStep::Compute {
expert: 3,
buffer: 1
},
]
);
}
#[test]
fn plan_double_buffer_overlap_and_raw_invariants_hold() {
for num_experts in 1..=8 {
let steps = plan_double_buffer(num_experts);
let pos = |pred: &dyn Fn(&PrefetchStep) -> bool| steps.iter().position(pred);
for n in 0..num_experts {
let compute_n =
pos(&|s| matches!(s, PrefetchStep::Compute { expert, .. } if *expert == n))
.expect("compute step present");
let await_n = pos(&|s| matches!(s, PrefetchStep::Await { expert } if *expert == n))
.expect("await step present");
assert!(await_n < compute_n, "await(n) must precede compute(n)");
if n + 1 < num_experts {
let prefetch_next = pos(
&|s| matches!(s, PrefetchStep::Prefetch { expert, .. } if *expert == n + 1),
)
.expect("prefetch(n+1) present");
assert!(
prefetch_next < compute_n,
"prefetch(n+1) must be issued before compute(n) to overlap"
);
}
let compute_slot = steps.iter().find_map(|s| match s {
PrefetchStep::Compute { expert, buffer } if *expert == n => Some(*buffer),
_ => None,
});
assert_eq!(compute_slot, Some(n % 2));
}
}
}
struct RecordingEp {
cpu_device: DeviceId,
log: Mutex<Vec<String>>,
next_fence_id: std::sync::atomic::AtomicU64,
}
impl RecordingEp {
fn new() -> Self {
Self {
cpu_device: DeviceId::cpu(),
log: Mutex::new(Vec::new()),
next_fence_id: std::sync::atomic::AtomicU64::new(1),
}
}
fn log(&self) -> Vec<String> {
self.log.lock().unwrap().clone()
}
}
impl ExecutionProvider for RecordingEp {
fn name(&self) -> &str {
"recording_ep"
}
fn device_type(&self) -> DeviceType {
self.cpu_device.device_type
}
fn device_id(&self) -> DeviceId {
self.cpu_device
}
fn initialize(&mut self, _config: &EpConfig) -> Result<()> {
Ok(())
}
fn shutdown(&mut self) -> Result<()> {
Ok(())
}
fn supports_op(
&self,
_op: &Node,
_opset: u64,
_shapes: &[Shape],
_input_dtypes: &[DataType],
_layouts: &[TensorLayout],
) -> KernelMatch {
KernelMatch::unsupported("recording_ep runs no kernels")
}
fn get_kernel(
&self,
_op: &Node,
_shapes: &[Vec<usize>],
_opset: u64,
) -> Result<Box<dyn onnx_runtime_ep_api::Kernel>> {
Err(onnx_runtime_ep_api::EpError::KernelFailed(
"recording_ep runs no kernels".into(),
))
}
fn allocate(&self, size: usize, alignment: usize) -> Result<DeviceBuffer> {
let layout = std::alloc::Layout::from_size_align(size.max(1), alignment)
.map_err(|_| onnx_runtime_ep_api::EpError::AlignmentError)?;
let ptr = unsafe { std::alloc::alloc(layout) };
if ptr.is_null() {
return Err(onnx_runtime_ep_api::EpError::OutOfMemory {
requested: size,
available: 0,
});
}
Ok(unsafe {
DeviceBuffer::from_raw_parts(ptr.cast(), self.cpu_device, size, alignment)
})
}
fn deallocate(&self, buffer: DeviceBuffer) -> Result<()> {
let size = buffer.len();
let alignment = buffer.alignment();
let ptr = buffer.into_raw().cast::<u8>();
let layout = std::alloc::Layout::from_size_align(size.max(1), alignment)
.expect("recording_ep allocated this layout");
unsafe { std::alloc::dealloc(ptr, layout) };
Ok(())
}
fn copy(&self, src: &DeviceBuffer, dst: &mut DeviceBuffer, size: usize) -> Result<()> {
if size != 0 {
unsafe {
std::ptr::copy_nonoverlapping(
src.as_ptr().cast::<u8>(),
dst.as_mut_ptr().cast::<u8>(),
size,
)
};
}
Ok(())
}
fn copy_async(
&self,
src: &DeviceBuffer,
dst: &mut DeviceBuffer,
size: usize,
) -> Result<Fence> {
self.copy(src, dst, size)?;
let id = self
.next_fence_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.log.lock().unwrap().push(format!("copy_async({id})"));
Ok(Fence::new(id))
}
fn wait_fence(&self, fence: &Fence) -> Result<()> {
self.log
.lock()
.unwrap()
.push(format!("wait_fence({})", fence.id));
Ok(())
}
fn sync(&self) -> Result<()> {
Ok(())
}
}
#[test]
fn drive_double_buffer_overlaps_and_orders_transfers() {
let mut ep = RecordingEp::new();
ep.initialize(&EpConfig::default()).unwrap();
let payloads: [[u8; 8]; 3] = [
[1, 1, 1, 1, 1, 1, 1, 1],
[2, 2, 2, 2, 2, 2, 2, 2],
[3, 3, 3, 3, 3, 3, 3, 3],
];
let sizes = [8usize, 8, 8];
let mut sources: Vec<DeviceBuffer> = Vec::new();
for p in &payloads {
let mut b = ep.allocate(8, 8).unwrap();
unsafe {
std::ptr::copy_nonoverlapping(p.as_ptr(), b.as_mut_ptr().cast::<u8>(), 8);
}
sources.push(b);
}
let mut buffers = [ep.allocate(8, 8).unwrap(), ep.allocate(8, 8).unwrap()];
let observed: Mutex<Vec<[u8; 8]>> = Mutex::new(Vec::new());
drive_double_buffer(&ep, &mut buffers, &sources, &sizes, |expert, weights| {
let mut got = [0u8; 8];
unsafe {
std::ptr::copy_nonoverlapping(weights.as_ptr().cast::<u8>(), got.as_mut_ptr(), 8);
}
let _ = expert;
observed.lock().unwrap().push(got);
Ok(())
})
.unwrap();
assert_eq!(
observed.into_inner().unwrap(),
vec![payloads[0], payloads[1], payloads[2]]
);
let log = ep.log();
assert_eq!(
log,
vec![
"copy_async(1)".to_string(),
"copy_async(2)".to_string(),
"wait_fence(1)".to_string(),
"copy_async(3)".to_string(),
"wait_fence(2)".to_string(),
"wait_fence(3)".to_string(),
]
);
let idx = |needle: &str| log.iter().position(|s| s == needle).unwrap();
assert!(
idx("copy_async(2)") < idx("wait_fence(1)"),
"expert 1 prefetch must be issued before awaiting expert 0"
);
assert!(
idx("copy_async(3)") < idx("wait_fence(2)"),
"expert 2 prefetch must be issued before awaiting expert 1"
);
for b in sources {
ep.deallocate(b).unwrap();
}
for b in buffers {
ep.deallocate(b).unwrap();
}
}
}