use vyre_driver::megakernel_execution::{
plan_megakernel_execution, select_megakernel_topology, select_megakernel_topology_stable,
MegakernelDeviceCapabilities, MegakernelExecutionPlan, MegakernelExecutionSample,
MegakernelExecutionTopology, MegakernelGraphShape, MegakernelMemoryBudget,
MegakernelMemoryError, MegakernelTopologyDecision,
};
use vyre_self_substrate::megakernel_schedule::{
try_schedule_via_scale_aware_samples_into, MegakernelScaleSample, MegakernelScheduleError,
};
use crate::backend::CudaTelemetrySnapshot;
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct CudaMegakernelScheduleSample {
pub dispatch_cost_ns: f64,
pub frontier_density: f64,
pub readback_bytes: u64,
}
impl CudaMegakernelScheduleSample {
#[must_use]
pub fn from_telemetry_snapshot(snapshot: CudaTelemetrySnapshot, dispatch_cost_ns: f64) -> Self {
let frontier_density = f64::from(snapshot.logical_thread_utilization_bps) / 10_000.0;
Self {
dispatch_cost_ns,
frontier_density,
readback_bytes: snapshot.readback_bytes,
}
}
fn execution_sample(self) -> MegakernelExecutionSample {
MegakernelExecutionSample {
dispatch_cost_ns: self.dispatch_cost_ns,
frontier_density: self.frontier_density,
readback_bytes: self.readback_bytes,
}
}
}
const CALLER_GATED: MegakernelDeviceCapabilities = MegakernelDeviceCapabilities::FUSION_CAPABLE;
#[must_use]
pub fn select_cuda_megakernel_topology(
sample: CudaMegakernelScheduleSample,
graph: MegakernelGraphShape,
memory: MegakernelMemoryBudget,
launch_overhead_ns: f64,
fusion_pressure: f64,
) -> MegakernelTopologyDecision {
select_megakernel_topology(
sample.execution_sample(),
graph,
memory,
launch_overhead_ns,
fusion_pressure,
CALLER_GATED,
)
}
#[must_use]
pub fn select_cuda_megakernel_topology_stable(
sample: CudaMegakernelScheduleSample,
graph: MegakernelGraphShape,
memory: MegakernelMemoryBudget,
launch_overhead_ns: f64,
fusion_pressure: f64,
previous_topology: MegakernelExecutionTopology,
) -> MegakernelTopologyDecision {
select_megakernel_topology_stable(
sample.execution_sample(),
graph,
memory,
launch_overhead_ns,
fusion_pressure,
previous_topology,
CALLER_GATED,
)
}
#[allow(clippy::too_many_arguments)]
pub fn plan_cuda_megakernel_execution(
sample: CudaMegakernelScheduleSample,
graph: MegakernelGraphShape,
bytes_per_node: u64,
bytes_per_edge: u64,
frontier_bytes: u64,
scratch_bytes: u64,
output_bytes: u64,
budget_bytes: u64,
launch_overhead_ns: f64,
fusion_pressure: f64,
) -> Result<MegakernelExecutionPlan, MegakernelMemoryError> {
plan_megakernel_execution(
sample.execution_sample(),
graph,
bytes_per_node,
bytes_per_edge,
frontier_bytes,
scratch_bytes,
output_bytes,
budget_bytes,
launch_overhead_ns,
fusion_pressure,
CALLER_GATED,
)
}
impl MegakernelScaleSample for CudaMegakernelScheduleSample {
fn dispatch_cost_ns(&self) -> f64 {
self.dispatch_cost_ns
}
fn frontier_density(&self) -> f64 {
self.frontier_density
}
fn readback_bytes(&self) -> u64 {
self.readback_bytes
}
}
pub fn schedule_megakernel_from_cuda_samples(
samples: &[CudaMegakernelScheduleSample],
launch_overhead_ns: f64,
n_steps: u32,
dt: f64,
) -> Result<Vec<f64>, MegakernelScheduleError> {
let mut out = Vec::new();
schedule_megakernel_from_cuda_samples_into(samples, launch_overhead_ns, n_steps, dt, &mut out)?;
Ok(out)
}
pub fn schedule_megakernel_from_cuda_samples_into(
samples: &[CudaMegakernelScheduleSample],
launch_overhead_ns: f64,
n_steps: u32,
dt: f64,
out: &mut Vec<f64>,
) -> Result<(), MegakernelScheduleError> {
try_schedule_via_scale_aware_samples_into(samples, launch_overhead_ns, n_steps, dt, out)
}
#[cfg(test)]
mod tests {
use super::{
schedule_megakernel_from_cuda_samples, schedule_megakernel_from_cuda_samples_into,
CudaMegakernelScheduleSample,
};
use crate::backend::CudaTelemetrySnapshot;
use vyre_self_substrate::megakernel_schedule::MegakernelScheduleError;
#[test]
fn telemetry_snapshot_maps_onto_a_scheduler_sample() {
let sample = CudaMegakernelScheduleSample::from_telemetry_snapshot(
CudaTelemetrySnapshot {
readback_bytes: 4096,
logical_thread_utilization_bps: 3750,
..CudaTelemetrySnapshot::default()
},
123.0,
);
assert_eq!(
sample,
CudaMegakernelScheduleSample {
dispatch_cost_ns: 123.0,
frontier_density: 0.375,
readback_bytes: 4096,
}
);
}
#[test]
fn scheduling_reuses_caller_owned_output_capacity() {
let samples = [
CudaMegakernelScheduleSample {
dispatch_cost_ns: 10.0,
frontier_density: 0.0,
readback_bytes: 0,
},
CudaMegakernelScheduleSample {
dispatch_cost_ns: 20.0,
frontier_density: 1.0,
readback_bytes: 4096,
},
];
let mut out = Vec::with_capacity(4);
let ptr = out.as_ptr();
schedule_megakernel_from_cuda_samples_into(&samples, 5.0, 8, 0.25, &mut out)
.expect("Fix: valid CUDA scheduler samples must schedule");
assert_eq!(out.len(), 2);
assert_eq!(out.as_ptr(), ptr);
assert!(out[1] > out[0]);
}
#[test]
fn scheduling_preserves_sample_validation_errors() {
let samples = [CudaMegakernelScheduleSample {
dispatch_cost_ns: 10.0,
frontier_density: 1.5,
readback_bytes: 0,
}];
let error = schedule_megakernel_from_cuda_samples(&samples, 0.0, 8, 0.25)
.expect_err("invalid frontier density must be rejected");
assert!(matches!(
error,
MegakernelScheduleError::InvalidFrontierDensity { index: 0, .. }
));
}
}