use super::TransferContext;
use super::{PhysicalLayout, TransferStrategy};
use crate::BlockId;
use crate::transfer::context::TransferCompleteNotification;
use crate::transfer::{can_use_whole_block_transfer, validate_layout_compatibility};
use anyhow::{Result, anyhow};
use cudarc::driver::{CudaStream, result as cuda_result};
use cudarc::runtime::sys::cudaStream_t;
use dynamo_memory::CudaMemPool;
use kvbm_kernels::MemcpyBatchMode;
use std::ffi::c_void;
use std::ops::Range;
use std::sync::Arc;
#[allow(clippy::too_many_arguments)]
pub fn execute_cuda_transfer(
src: &PhysicalLayout,
dst: &PhysicalLayout,
src_block_ids: &[BlockId],
dst_block_ids: &[BlockId],
layer_range: Option<Range<usize>>,
strategy: TransferStrategy,
cuda_stream: Option<Arc<CudaStream>>,
ctx: &TransferContext,
) -> Result<TransferCompleteNotification> {
let src_layout = src.layout();
let dst_layout = dst.layout();
if src_layout.num_layers() != dst_layout.num_layers() {
return Err(anyhow!(
"Layouts have incompatible layer counts: src={}, dst={}",
src_layout.num_layers(),
dst_layout.num_layers()
));
}
if src_layout.outer_dim() != dst_layout.outer_dim() {
return Err(anyhow!(
"Layouts have incompatible outer dimensions: src={}, dst={}",
src_layout.outer_dim(),
dst_layout.outer_dim()
));
}
validate_layout_compatibility(src, dst)?;
let layers = layer_range.clone().unwrap_or(0..src_layout.num_layers());
let use_whole_block = can_use_whole_block_transfer(src, dst, layer_range.as_ref());
let caller_manages_sync = cuda_stream.is_some();
let stream = if let Some(s) = cuda_stream {
s
} else {
match strategy {
TransferStrategy::CudaAsyncD2H => ctx.next_d2h_streams(),
_ => ctx.next_h2d_streams(), }
};
let strategy_name = match strategy {
TransferStrategy::CudaAsyncH2D => "H2D",
TransferStrategy::CudaAsyncD2H => "D2H",
TransferStrategy::CudaAsyncD2D => "D2D",
_ => "Unknown",
};
match strategy {
TransferStrategy::CudaAsyncH2D
| TransferStrategy::CudaAsyncD2H
| TransferStrategy::CudaAsyncD2D => {
if use_whole_block {
tracing::debug!(
strategy = strategy_name,
num_blocks = src_block_ids.len(),
bytes_per_block = src_layout.config().bytes_per_block(),
"Using whole-block transfer (auto direction)"
);
execute_whole_block_cuda(src, dst, src_block_ids, dst_block_ids, stream.as_ref())?;
} else {
tracing::debug!(
strategy = strategy_name,
num_blocks = src_block_ids.len(),
num_layers = layers.len(),
"Using vectorized_copy for FC↔LW transfer"
);
execute_fc_lw_vectorized(
src,
dst,
src_block_ids,
dst_block_ids,
layers.clone(),
stream.as_ref(),
ctx.cuda_pool(),
)?;
}
}
_ => {
return Err(anyhow!("Invalid CUDA transfer strategy: {:?}", strategy));
}
}
if caller_manages_sync {
return Ok(TransferCompleteNotification::completed());
}
if matches!(
strategy,
TransferStrategy::CudaAsyncH2D
| TransferStrategy::CudaAsyncD2H
| TransferStrategy::CudaAsyncD2D
) {
let event = stream.record_event(None)?;
Ok(ctx.register_cuda_event(event))
} else {
Ok(TransferCompleteNotification::completed())
}
}
fn execute_whole_block_cuda(
src: &PhysicalLayout,
dst: &PhysicalLayout,
src_block_ids: &[BlockId],
dst_block_ids: &[BlockId],
stream: &cudarc::driver::CudaStream,
) -> Result<()> {
let bytes_per_block = src.layout().config().bytes_per_block();
let num_blocks = src_block_ids.len();
if num_blocks == 0 {
return Ok(());
}
let mut src_ptrs: Vec<*const std::ffi::c_void> = Vec::with_capacity(num_blocks);
let mut dst_ptrs: Vec<*mut std::ffi::c_void> = Vec::with_capacity(num_blocks);
for (&src_block_id, &dst_block_id) in src_block_ids.iter().zip(dst_block_ids.iter()) {
let src_region = src.memory_region(src_block_id, 0, 0)?;
let dst_region = dst.memory_region(dst_block_id, 0, 0)?;
src_ptrs.push(src_region.addr() as *const std::ffi::c_void);
dst_ptrs.push(dst_region.addr() as *mut std::ffi::c_void);
}
let status = unsafe {
kvbm_kernels::memcpy_batch(
src_ptrs.as_ptr(),
dst_ptrs.as_ptr(),
bytes_per_block,
num_blocks,
MemcpyBatchMode::BatchedWithFallback,
stream.cu_stream() as cudarc::runtime::sys::cudaStream_t,
)
};
if status != cudarc::runtime::sys::cudaError::cudaSuccess {
return Err(anyhow!("memcpy_batch failed: {:?}", status));
}
tracing::debug!(
num_blocks,
bytes_per_block,
batch_available = kvbm_kernels::is_memcpy_batch_available(),
"Whole-block transfer completed"
);
Ok(())
}
fn execute_fc_lw_vectorized(
src: &PhysicalLayout,
dst: &PhysicalLayout,
src_block_ids: &[BlockId],
dst_block_ids: &[BlockId],
layers: Range<usize>,
stream: &CudaStream,
pool: &CudaMemPool,
) -> Result<()> {
stream.context().bind_to_thread()?;
let src_layout = src.layout();
let nl = layers.len();
let no = src_layout.outer_dim();
let chunk_size =
src_layout.page_size() * src_layout.inner_dim() * src_layout.dtype_width_bytes();
let num_blocks = src_block_ids.len();
let total_chunks = num_blocks * nl * no;
if total_chunks == 0 {
return Ok(());
}
let mut src_ptrs: Vec<usize> = Vec::with_capacity(total_chunks);
let mut dst_ptrs: Vec<usize> = Vec::with_capacity(total_chunks);
for (&src_block_id, &dst_block_id) in src_block_ids.iter().zip(dst_block_ids.iter()) {
for layer_id in layers.clone() {
for outer_id in 0..no {
let src_region = src.memory_region(src_block_id, layer_id, outer_id)?;
let dst_region = dst.memory_region(dst_block_id, layer_id, outer_id)?;
src_ptrs.push(src_region.addr());
dst_ptrs.push(dst_region.addr());
}
}
}
let src_ptrs_device = pool.alloc_async(total_chunks * std::mem::size_of::<usize>(), stream)?;
let dst_ptrs_device = pool.alloc_async(total_chunks * std::mem::size_of::<usize>(), stream)?;
unsafe {
cuda_result::memcpy_htod_async(
src_ptrs_device,
std::slice::from_raw_parts(
src_ptrs.as_ptr() as *const u8,
total_chunks * std::mem::size_of::<usize>(),
),
stream.cu_stream(),
)?;
cuda_result::memcpy_htod_async(
dst_ptrs_device,
std::slice::from_raw_parts(
dst_ptrs.as_ptr() as *const u8,
total_chunks * std::mem::size_of::<usize>(),
),
stream.cu_stream(),
)?;
}
let pointers_transfered_event = stream.record_event(None)?;
let status = unsafe {
kvbm_kernels::vectorized_copy(
src_ptrs_device as *mut *mut c_void,
dst_ptrs_device as *mut *mut c_void,
chunk_size,
total_chunks as i32,
stream.cu_stream() as cudaStream_t,
)
};
pool.free_async(src_ptrs_device, stream)?;
pool.free_async(dst_ptrs_device, stream)?;
if status != cudarc::runtime::sys::cudaError::cudaSuccess {
return Err(anyhow!("vectorized_copy failed: {:?}", status));
}
tracing::debug!(
total_chunks,
chunk_size,
"FC↔LW vectorized_copy transfer completed"
);
pointers_transfered_event.synchronize()?;
Ok(())
}