use crate::error::{OptimError, Result};
use std::collections::VecDeque;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PipelineSchedule {
GPipe,
OneForwardOneBackward,
}
impl PipelineSchedule {
pub fn name(self) -> &'static str {
match self {
PipelineSchedule::GPipe => "GPipe",
PipelineSchedule::OneForwardOneBackward => "1F1B",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OpKind {
Forward,
Backward,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StageCost {
pub forward: f64,
pub backward: f64,
}
impl StageCost {
pub fn new(forward: f64, backward: f64) -> Self {
Self { forward, backward }
}
pub fn uniform(value: f64) -> Self {
Self {
forward: value,
backward: value,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct PipelineOp {
pub stage: usize,
pub micro_batch: usize,
pub kind: OpKind,
pub start: f64,
pub end: f64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PipelineConfig {
pub num_stages: usize,
pub num_micro_batches: usize,
}
impl PipelineConfig {
pub fn new(num_stages: usize, num_micro_batches: usize) -> Result<Self> {
if num_stages == 0 {
return Err(OptimError::InvalidConfig(
"num_stages must be at least 1".to_string(),
));
}
if num_micro_batches == 0 {
return Err(OptimError::InvalidConfig(
"num_micro_batches must be at least 1".to_string(),
));
}
Ok(Self {
num_stages,
num_micro_batches,
})
}
pub fn analytical_bubble_fraction(&self) -> f64 {
let p = self.num_stages as f64;
let m = self.num_micro_batches as f64;
(p - 1.0) / (m + p - 1.0)
}
pub fn analytical_utilization(&self) -> f64 {
let p = self.num_stages as f64;
let m = self.num_micro_batches as f64;
m / (m + p - 1.0)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct PipelineMetrics {
pub bubble_fraction: f64,
pub utilization: f64,
pub peak_activation_stash: usize,
pub per_stage_peak_stash: Vec<usize>,
pub makespan: f64,
pub throughput: f64,
}
#[derive(Debug, Clone)]
pub struct PipelineExecution {
pub schedule: PipelineSchedule,
pub config: PipelineConfig,
pub ops: Vec<PipelineOp>,
pub metrics: PipelineMetrics,
}
impl PipelineExecution {
pub fn ops(&self) -> &[PipelineOp] {
&self.ops
}
pub fn metrics(&self) -> &PipelineMetrics {
&self.metrics
}
pub fn makespan(&self) -> f64 {
self.metrics.makespan
}
pub fn total_busy_time(&self) -> f64 {
self.ops.iter().map(|op| op.end - op.start).sum()
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct StageRange {
pub stage: usize,
pub start_layer: usize,
pub end_layer: usize,
pub load: f64,
}
impl StageRange {
pub fn num_layers(&self) -> usize {
self.end_layer - self.start_layer
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct StagePartitioner;
impl StagePartitioner {
pub fn new() -> Self {
Self
}
pub fn partition(&self, layer_costs: &[f64], num_stages: usize) -> Result<Vec<StageRange>> {
let num_layers = layer_costs.len();
if num_stages == 0 {
return Err(OptimError::InvalidConfig(
"num_stages must be at least 1".to_string(),
));
}
if num_layers == 0 {
return Err(OptimError::InvalidConfig(
"layer_costs must not be empty".to_string(),
));
}
if num_stages > num_layers {
return Err(OptimError::InvalidConfig(format!(
"num_stages {num_stages} exceeds number of layers {num_layers}: \
cannot form non-empty contiguous stages"
)));
}
for (i, &cost) in layer_costs.iter().enumerate() {
if !cost.is_finite() || cost < 0.0 {
return Err(OptimError::InvalidConfig(format!(
"layer {i} cost {cost} must be finite and non-negative"
)));
}
}
let mut prefix = vec![0.0f64; num_layers + 1];
for i in 0..num_layers {
prefix[i + 1] = prefix[i] + layer_costs[i];
}
let mut dp = vec![vec![f64::INFINITY; num_layers + 1]; num_stages + 1];
let mut choice = vec![vec![0usize; num_layers + 1]; num_stages + 1];
dp[1][1..=num_layers].copy_from_slice(&prefix[1..=num_layers]);
for stages in 2..=num_stages {
for i in stages..=num_layers {
let mut best = f64::INFINITY;
let mut best_j = stages - 1;
for j in (stages - 1)..i {
let last_load = prefix[i] - prefix[j];
let candidate = dp[stages - 1][j].max(last_load);
if candidate < best {
best = candidate;
best_j = j;
}
}
dp[stages][i] = best;
choice[stages][i] = best_j;
}
}
let mut ranges: Vec<StageRange> = Vec::with_capacity(num_stages);
let mut end = num_layers;
let mut stages = num_stages;
while stages >= 1 {
let start = if stages == 1 { 0 } else { choice[stages][end] };
ranges.push(StageRange {
stage: stages - 1,
start_layer: start,
end_layer: end,
load: prefix[end] - prefix[start],
});
end = start;
stages -= 1;
}
ranges.reverse();
Ok(ranges)
}
pub fn optimal_max_load(&self, layer_costs: &[f64], num_stages: usize) -> Result<f64> {
let ranges = self.partition(layer_costs, num_stages)?;
Ok(ranges.iter().map(|range| range.load).fold(0.0f64, f64::max))
}
}
#[inline]
fn op_index(stage: usize, micro: usize, kind: OpKind, num_micro: usize) -> usize {
let kind_idx = match kind {
OpKind::Forward => 0,
OpKind::Backward => 1,
};
(stage * num_micro + micro) * 2 + kind_idx
}
#[inline]
fn decode_index(index: usize, num_micro: usize) -> (usize, usize, OpKind) {
let kind = if index.is_multiple_of(2) {
OpKind::Forward
} else {
OpKind::Backward
};
let rest = index / 2;
let micro = rest % num_micro;
let stage = rest / num_micro;
(stage, micro, kind)
}
fn gpipe_stage_orders(num_stages: usize, num_micro: usize) -> Vec<Vec<(usize, OpKind)>> {
let mut orders = Vec::with_capacity(num_stages);
for _ in 0..num_stages {
let mut order = Vec::with_capacity(2 * num_micro);
for micro in 0..num_micro {
order.push((micro, OpKind::Forward));
}
for micro in (0..num_micro).rev() {
order.push((micro, OpKind::Backward));
}
orders.push(order);
}
orders
}
fn one_f_one_b_stage_orders(num_stages: usize, num_micro: usize) -> Vec<Vec<(usize, OpKind)>> {
let mut orders = Vec::with_capacity(num_stages);
for stage in 0..num_stages {
let warmup = (num_stages - 1 - stage).min(num_micro);
let steady = num_micro - warmup;
let mut order = Vec::with_capacity(2 * num_micro);
for micro in 0..warmup {
order.push((micro, OpKind::Forward));
}
for k in 0..steady {
order.push((warmup + k, OpKind::Forward));
order.push((k, OpKind::Backward));
}
for micro in steady..num_micro {
order.push((micro, OpKind::Backward));
}
orders.push(order);
}
orders
}
fn compute_timeline(
num_stages: usize,
num_micro: usize,
stage_orders: &[Vec<(usize, OpKind)>],
stage_costs: &[StageCost],
) -> Result<(Vec<PipelineOp>, f64)> {
let num_ops = num_stages * num_micro * 2;
let mut preds: Vec<Vec<usize>> = vec![Vec::new(); num_ops];
for (stage, order) in stage_orders.iter().enumerate() {
for window in order.windows(2) {
let prev = op_index(stage, window[0].0, window[0].1, num_micro);
let cur = op_index(stage, window[1].0, window[1].1, num_micro);
preds[cur].push(prev);
}
}
for micro in 0..num_micro {
for stage in 0..num_stages {
let forward = op_index(stage, micro, OpKind::Forward, num_micro);
if stage > 0 {
preds[forward].push(op_index(stage - 1, micro, OpKind::Forward, num_micro));
}
let backward = op_index(stage, micro, OpKind::Backward, num_micro);
if stage + 1 < num_stages {
preds[backward].push(op_index(stage + 1, micro, OpKind::Backward, num_micro));
}
preds[backward].push(forward);
}
}
let mut indeg = vec![0usize; num_ops];
let mut succ: Vec<Vec<usize>> = vec![Vec::new(); num_ops];
for (op, plist) in preds.iter().enumerate() {
indeg[op] = plist.len();
for &pred in plist {
succ[pred].push(op);
}
}
let mut start = vec![0.0f64; num_ops];
let mut end = vec![0.0f64; num_ops];
let mut queue: VecDeque<usize> = VecDeque::new();
for (op, °) in indeg.iter().enumerate() {
if deg == 0 {
queue.push_back(op);
}
}
let mut processed = 0usize;
while let Some(op) = queue.pop_front() {
let mut earliest = 0.0f64;
for &pred in &preds[op] {
if end[pred] > earliest {
earliest = end[pred];
}
}
let (stage, _micro, kind) = decode_index(op, num_micro);
let cost = match kind {
OpKind::Forward => stage_costs[stage].forward,
OpKind::Backward => stage_costs[stage].backward,
};
start[op] = earliest;
end[op] = earliest + cost;
processed += 1;
for &next in &succ[op] {
indeg[next] -= 1;
if indeg[next] == 0 {
queue.push_back(next);
}
}
}
if processed != num_ops {
return Err(OptimError::InvalidState(
"pipeline dependency graph is cyclic; schedule is infeasible".to_string(),
));
}
let mut makespan = 0.0f64;
let mut ops = Vec::with_capacity(num_ops);
for op in 0..num_ops {
let (stage, micro, kind) = decode_index(op, num_micro);
if end[op] > makespan {
makespan = end[op];
}
ops.push(PipelineOp {
stage,
micro_batch: micro,
kind,
start: start[op],
end: end[op],
});
}
ops.sort_by(|a, b| {
a.start
.partial_cmp(&b.start)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.stage.cmp(&b.stage))
.then((a.kind as usize).cmp(&(b.kind as usize)))
.then(a.micro_batch.cmp(&b.micro_batch))
});
Ok((ops, makespan))
}
fn compute_peak_stash(stage_orders: &[Vec<(usize, OpKind)>]) -> Vec<usize> {
let mut peaks = Vec::with_capacity(stage_orders.len());
for order in stage_orders {
let mut current = 0i64;
let mut peak = 0i64;
for &(_, kind) in order {
match kind {
OpKind::Forward => {
current += 1;
if current > peak {
peak = current;
}
}
OpKind::Backward => {
current -= 1;
}
}
}
peaks.push(peak.max(0) as usize);
}
peaks
}
#[derive(Debug, Clone, Copy)]
pub struct PipelineScheduler {
config: PipelineConfig,
}
impl PipelineScheduler {
pub fn new(config: PipelineConfig) -> Self {
Self { config }
}
pub fn config(&self) -> &PipelineConfig {
&self.config
}
pub fn schedule(
&self,
schedule_kind: PipelineSchedule,
stage_costs: &[StageCost],
) -> Result<PipelineExecution> {
let num_stages = self.config.num_stages;
let num_micro = self.config.num_micro_batches;
if stage_costs.len() != num_stages {
return Err(OptimError::DimensionMismatch(format!(
"expected {num_stages} stage costs (one per stage), got {}",
stage_costs.len()
)));
}
for (stage, cost) in stage_costs.iter().enumerate() {
if !cost.forward.is_finite() || cost.forward <= 0.0 {
return Err(OptimError::InvalidConfig(format!(
"stage {stage} forward cost {} must be finite and positive",
cost.forward
)));
}
if !cost.backward.is_finite() || cost.backward <= 0.0 {
return Err(OptimError::InvalidConfig(format!(
"stage {stage} backward cost {} must be finite and positive",
cost.backward
)));
}
}
let stage_orders = match schedule_kind {
PipelineSchedule::GPipe => gpipe_stage_orders(num_stages, num_micro),
PipelineSchedule::OneForwardOneBackward => {
one_f_one_b_stage_orders(num_stages, num_micro)
}
};
let (ops, makespan) = compute_timeline(num_stages, num_micro, &stage_orders, stage_costs)?;
let per_stage_peak_stash = compute_peak_stash(&stage_orders);
let peak_activation_stash = per_stage_peak_stash.iter().copied().max().unwrap_or(0);
let total_busy: f64 = stage_costs
.iter()
.map(|cost| (cost.forward + cost.backward) * num_micro as f64)
.sum();
let capacity = num_stages as f64 * makespan;
let utilization = if capacity > 0.0 {
(total_busy / capacity).min(1.0)
} else {
0.0
};
let bubble_fraction = (1.0 - utilization).max(0.0);
let throughput = if makespan > 0.0 {
num_micro as f64 / makespan
} else {
0.0
};
let metrics = PipelineMetrics {
bubble_fraction,
utilization,
peak_activation_stash,
per_stage_peak_stash,
makespan,
throughput,
};
Ok(PipelineExecution {
schedule: schedule_kind,
config: self.config,
ops,
metrics,
})
}
pub fn schedule_uniform(
&self,
schedule_kind: PipelineSchedule,
forward: f64,
backward: f64,
) -> Result<PipelineExecution> {
let stage_costs = vec![StageCost::new(forward, backward); self.config.num_stages];
self.schedule(schedule_kind, &stage_costs)
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
fn brute_force_max_load(layer_costs: &[f64], num_stages: usize) -> f64 {
let num_layers = layer_costs.len();
let mut prefix = vec![0.0f64; num_layers + 1];
for i in 0..num_layers {
prefix[i + 1] = prefix[i] + layer_costs[i];
}
fn rec(prefix: &[f64], start: usize, stages: usize, num_layers: usize) -> f64 {
if stages == 1 {
return prefix[num_layers] - prefix[start];
}
let mut best = f64::INFINITY;
let last_end = num_layers - (stages - 1);
for end in (start + 1)..=last_end {
let first = prefix[end] - prefix[start];
let rest = rec(prefix, end, stages - 1, num_layers);
let candidate = first.max(rest);
if candidate < best {
best = candidate;
}
}
best
}
rec(&prefix, 0, num_stages, num_layers)
}
fn assert_contiguous_cover(ranges: &[StageRange], num_layers: usize, num_stages: usize) {
assert_eq!(ranges.len(), num_stages, "wrong number of stages");
assert_eq!(
ranges[0].start_layer, 0,
"first stage must start at layer 0"
);
assert_eq!(
ranges[num_stages - 1].end_layer,
num_layers,
"last stage must end at the final layer"
);
for (i, range) in ranges.iter().enumerate() {
assert_eq!(range.stage, i, "stage index out of order");
assert!(range.num_layers() >= 1, "every stage must be non-empty");
if i + 1 < ranges.len() {
assert_eq!(
range.end_layer,
ranges[i + 1].start_layer,
"stages must be contiguous"
);
}
}
}
#[test]
fn test_partition_balances_uniform_load() {
let partitioner = StagePartitioner::new();
let costs = vec![1.0f64; 8];
let ranges = partitioner.partition(&costs, 4).unwrap();
assert_contiguous_cover(&ranges, 8, 4);
for range in &ranges {
assert_eq!(range.num_layers(), 2);
assert_relative_eq!(range.load, 2.0, epsilon = 1e-12);
}
let max_load = ranges.iter().map(|r| r.load).fold(0.0, f64::max);
assert_relative_eq!(max_load, 2.0, epsilon = 1e-12);
}
#[test]
fn test_partition_matches_brute_force_optimum() {
let partitioner = StagePartitioner::new();
let cases: &[(Vec<f64>, usize)] = &[
(vec![3.0, 1.0, 1.0, 1.0, 3.0, 1.0], 3),
(vec![5.0, 2.0, 4.0, 1.0, 1.0, 9.0, 3.0, 2.0], 4),
(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0], 2),
(vec![10.0, 1.0, 1.0, 1.0, 1.0], 5),
(vec![2.0, 2.0, 2.0, 2.0, 2.0, 2.0], 3),
];
for (costs, num_stages) in cases {
let ranges = partitioner.partition(costs, *num_stages).unwrap();
assert_contiguous_cover(&ranges, costs.len(), *num_stages);
let dp_max = ranges.iter().map(|r| r.load).fold(0.0, f64::max);
let optimum = brute_force_max_load(costs, *num_stages);
assert_relative_eq!(dp_max, optimum, epsilon = 1e-9);
let reported = partitioner.optimal_max_load(costs, *num_stages).unwrap();
assert_relative_eq!(reported, optimum, epsilon = 1e-9);
let naive = naive_equal_count_max_load(costs, *num_stages);
assert!(
dp_max <= naive + 1e-9,
"balanced split must beat naive split"
);
}
}
fn naive_equal_count_max_load(costs: &[f64], num_stages: usize) -> f64 {
let num_layers = costs.len();
let base = num_layers / num_stages;
let rem = num_layers % num_stages;
let mut idx = 0usize;
let mut max_load = 0.0f64;
for stage in 0..num_stages {
let count = if stage < rem { base + 1 } else { base };
let load: f64 = costs[idx..idx + count].iter().sum();
if load > max_load {
max_load = load;
}
idx += count;
}
max_load
}
#[test]
fn test_partition_invalid_configs() {
let partitioner = StagePartitioner::new();
assert!(partitioner.partition(&[1.0, 2.0], 0).is_err());
assert!(partitioner.partition(&[], 1).is_err());
assert!(partitioner.partition(&[1.0, 2.0], 3).is_err());
assert!(partitioner.partition(&[1.0, -1.0, 2.0], 2).is_err());
assert!(partitioner.partition(&[1.0, f64::NAN], 2).is_err());
assert!(partitioner.partition(&[1.0, 2.0, 3.0], 3).is_ok());
}
#[test]
fn test_pipeline_config_validation() {
assert!(PipelineConfig::new(0, 4).is_err());
assert!(PipelineConfig::new(4, 0).is_err());
assert!(PipelineConfig::new(1, 1).is_ok());
assert!(PipelineConfig::new(4, 8).is_ok());
}
#[test]
fn test_gpipe_bubble_fraction_matches_formula() {
let cases = [(2usize, 2usize), (4, 8), (4, 1), (8, 16), (3, 5), (1, 4)];
for (p, m) in cases {
let config = PipelineConfig::new(p, m).unwrap();
let scheduler = PipelineScheduler::new(config);
let exec = scheduler
.schedule_uniform(PipelineSchedule::GPipe, 1.0, 1.0)
.unwrap();
let analytical = config.analytical_bubble_fraction();
assert_relative_eq!(
analytical,
(p as f64 - 1.0) / (m as f64 + p as f64 - 1.0),
epsilon = 1e-12
);
assert_relative_eq!(exec.metrics.bubble_fraction, analytical, epsilon = 1e-9);
assert_relative_eq!(
exec.metrics.utilization,
config.analytical_utilization(),
epsilon = 1e-9
);
assert_relative_eq!(
exec.metrics.bubble_fraction + exec.metrics.utilization,
1.0,
epsilon = 1e-9
);
}
}
#[test]
fn test_gpipe_generated_idle_matches_analytical_bubble() {
let cases = [(2usize, 4usize), (4, 8), (3, 6), (5, 10)];
for (p, m) in cases {
let config = PipelineConfig::new(p, m).unwrap();
let scheduler = PipelineScheduler::new(config);
let exec = scheduler
.schedule_uniform(PipelineSchedule::GPipe, 1.0, 1.0)
.unwrap();
let busy: f64 = exec.ops.iter().map(|op| op.end - op.start).sum();
let capacity = p as f64 * exec.metrics.makespan;
let idle_fraction = 1.0 - busy / capacity;
assert_relative_eq!(
idle_fraction,
config.analytical_bubble_fraction(),
epsilon = 1e-9
);
assert_relative_eq!(
exec.metrics.makespan,
2.0 * (m as f64 + p as f64 - 1.0),
epsilon = 1e-9
);
}
}
#[test]
fn test_one_f_one_b_lower_activation_stash() {
let cases = [(4usize, 8usize), (8, 16), (4, 4), (3, 10), (6, 2)];
for (p, m) in cases {
let config = PipelineConfig::new(p, m).unwrap();
let scheduler = PipelineScheduler::new(config);
let gpipe = scheduler
.schedule_uniform(PipelineSchedule::GPipe, 1.0, 1.0)
.unwrap();
let one_f_one_b = scheduler
.schedule_uniform(PipelineSchedule::OneForwardOneBackward, 1.0, 1.0)
.unwrap();
assert_eq!(gpipe.metrics.peak_activation_stash, m);
assert_eq!(
one_f_one_b.metrics.peak_activation_stash,
p.min(m),
"1F1B peak stash should equal min(P, M)"
);
assert!(
one_f_one_b.metrics.peak_activation_stash <= gpipe.metrics.peak_activation_stash,
"1F1B peak must not exceed GPipe peak"
);
assert!(
one_f_one_b.metrics.peak_activation_stash <= p,
"1F1B peak must not exceed pipeline depth P"
);
for &stage_peak in &one_f_one_b.metrics.per_stage_peak_stash {
assert!(stage_peak <= p, "per-stage 1F1B stash must be <= P");
}
}
}
#[test]
fn test_one_f_one_b_strictly_lower_stash_for_large_m() {
let config = PipelineConfig::new(4, 16).unwrap();
let scheduler = PipelineScheduler::new(config);
let gpipe = scheduler
.schedule_uniform(PipelineSchedule::GPipe, 1.0, 1.0)
.unwrap();
let one_f_one_b = scheduler
.schedule_uniform(PipelineSchedule::OneForwardOneBackward, 1.0, 1.0)
.unwrap();
assert_eq!(gpipe.metrics.peak_activation_stash, 16);
assert_eq!(one_f_one_b.metrics.peak_activation_stash, 4);
assert!(one_f_one_b.metrics.peak_activation_stash < gpipe.metrics.peak_activation_stash);
}
#[test]
fn test_gpipe_and_one_f_one_b_same_bubble_and_makespan_uniform() {
let cases = [(2usize, 2usize), (4, 8), (3, 7), (5, 5)];
for (p, m) in cases {
let config = PipelineConfig::new(p, m).unwrap();
let scheduler = PipelineScheduler::new(config);
let gpipe = scheduler
.schedule_uniform(PipelineSchedule::GPipe, 1.0, 1.0)
.unwrap();
let one_f_one_b = scheduler
.schedule_uniform(PipelineSchedule::OneForwardOneBackward, 1.0, 1.0)
.unwrap();
assert_relative_eq!(
gpipe.metrics.makespan,
one_f_one_b.metrics.makespan,
epsilon = 1e-9
);
assert_relative_eq!(
gpipe.metrics.makespan,
2.0 * (m as f64 + p as f64 - 1.0),
epsilon = 1e-9
);
assert_relative_eq!(
gpipe.metrics.bubble_fraction,
one_f_one_b.metrics.bubble_fraction,
epsilon = 1e-9
);
}
}
#[test]
fn test_throughput_increases_with_micro_batches() {
for schedule in [
PipelineSchedule::GPipe,
PipelineSchedule::OneForwardOneBackward,
] {
let micro_batches = [1usize, 2, 4, 8, 16];
let mut previous = 0.0f64;
for &m in µ_batches {
let config = PipelineConfig::new(4, m).unwrap();
let scheduler = PipelineScheduler::new(config);
let exec = scheduler.schedule_uniform(schedule, 1.0, 1.0).unwrap();
assert!(
exec.metrics.throughput > previous,
"throughput must increase with M for {} (M={m})",
schedule.name()
);
previous = exec.metrics.throughput;
}
}
}
#[test]
fn test_schedule_structure_is_valid() {
let config = PipelineConfig::new(4, 6).unwrap();
let scheduler = PipelineScheduler::new(config);
for schedule in [
PipelineSchedule::GPipe,
PipelineSchedule::OneForwardOneBackward,
] {
let exec = scheduler.schedule_uniform(schedule, 1.0, 2.0).unwrap();
assert_eq!(exec.ops.len(), 4 * 6 * 2);
for stage in 0..4 {
for micro in 0..6 {
let forward = exec
.ops
.iter()
.find(|op| {
op.stage == stage
&& op.micro_batch == micro
&& op.kind == OpKind::Forward
})
.unwrap();
let backward = exec
.ops
.iter()
.find(|op| {
op.stage == stage
&& op.micro_batch == micro
&& op.kind == OpKind::Backward
})
.unwrap();
assert!(forward.end <= backward.start + 1e-9);
assert_relative_eq!(forward.end - forward.start, 1.0, epsilon = 1e-9);
assert_relative_eq!(backward.end - backward.start, 2.0, epsilon = 1e-9);
}
}
for micro in 0..6 {
for stage in 0..3 {
let here = exec
.ops
.iter()
.find(|op| {
op.stage == stage
&& op.micro_batch == micro
&& op.kind == OpKind::Forward
})
.unwrap();
let next = exec
.ops
.iter()
.find(|op| {
op.stage == stage + 1
&& op.micro_batch == micro
&& op.kind == OpKind::Forward
})
.unwrap();
assert!(here.end <= next.start + 1e-9);
}
}
}
}
#[test]
fn test_schedule_invalid_costs() {
let config = PipelineConfig::new(3, 4).unwrap();
let scheduler = PipelineScheduler::new(config);
let too_few = vec![StageCost::uniform(1.0); 2];
assert!(scheduler
.schedule(PipelineSchedule::GPipe, &too_few)
.is_err());
let bad = vec![
StageCost::new(1.0, 1.0),
StageCost::new(0.0, 1.0),
StageCost::new(1.0, 1.0),
];
assert!(scheduler.schedule(PipelineSchedule::GPipe, &bad).is_err());
let infinite = vec![
StageCost::new(1.0, 1.0),
StageCost::new(1.0, f64::INFINITY),
StageCost::new(1.0, 1.0),
];
assert!(scheduler
.schedule(PipelineSchedule::GPipe, &infinite)
.is_err());
}
#[test]
fn test_non_uniform_costs_bottleneck_dominates_makespan() {
let config = PipelineConfig::new(3, 8).unwrap();
let scheduler = PipelineScheduler::new(config);
let stage_costs = [
StageCost::new(1.0, 1.0),
StageCost::new(4.0, 4.0),
StageCost::new(1.0, 1.0),
];
let exec = scheduler
.schedule(PipelineSchedule::OneForwardOneBackward, &stage_costs)
.unwrap();
assert!(exec.metrics.makespan >= 64.0 - 1e-9);
assert!(exec.metrics.utilization > 0.0 && exec.metrics.utilization <= 1.0);
assert!(exec.metrics.bubble_fraction >= 0.0 && exec.metrics.bubble_fraction < 1.0);
}
}