Files
rustytorch/crates/training/rtx-distributed/src/pipeline_parallel/executor.rs
T
2026-03-04 00:08:42 +00:00

334 lines
12 KiB
Rust

//! Synchronous pipeline executor.
use crate::error::{DistributedError, Result};
use parking_lot::RwLock;
use rtx_tensor::Tensor;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use super::config::PipelineConfig;
use super::micro_batch::MicroBatch;
use super::scheduler::{PipelineOp, PipelineScheduler};
use super::stage::{PipelineStage, StageState};
use super::stage_module::StageModule;
use super::stats::{PipelineStats, StageStats};
/// Executes the pipeline schedule
pub struct PipelineExecutor {
/// Configuration
config: PipelineConfig,
/// Pipeline stages
stages: Vec<Arc<PipelineStage>>,
/// Scheduler
scheduler: PipelineScheduler,
/// Is pipeline running
running: AtomicBool,
/// Current step
current_step: AtomicUsize,
/// Statistics
stats: RwLock<PipelineStats>,
}
impl PipelineExecutor {
/// Create a new pipeline executor
pub fn new(config: PipelineConfig, device_ids: Vec<i32>) -> Result<Self> {
if device_ids.len() != config.num_stages {
return Err(DistributedError::configuration(format!(
"Device count ({}) must match stage count ({})",
device_ids.len(),
config.num_stages
)));
}
let stages: Vec<Arc<PipelineStage>> = device_ids
.iter()
.enumerate()
.map(|(i, &device_id)| Arc::new(PipelineStage::new(i, device_id)))
.collect();
let scheduler = PipelineScheduler::new(config.clone());
Ok(Self {
config,
stages,
scheduler,
running: AtomicBool::new(false),
current_step: AtomicUsize::new(0),
stats: RwLock::new(PipelineStats::default()),
})
}
/// Execute the pipeline for one mini-batch
pub fn execute(&self, input_batches: Vec<MicroBatch>) -> Result<Vec<MicroBatch>> {
if input_batches.len() != self.config.num_micro_batches {
return Err(DistributedError::configuration(format!(
"Input batch count ({}) must match micro-batch count ({})",
input_batches.len(),
self.config.num_micro_batches
)));
}
self.running.store(true, Ordering::SeqCst);
let start = Instant::now();
// Enqueue input batches to first stage
for batch in input_batches {
self.stages[0].enqueue_forward(batch);
}
// Execute schedule
let schedule = self.scheduler.schedule();
let mut outputs = Vec::new();
for (step_idx, step) in schedule.iter().enumerate() {
self.current_step.store(step_idx, Ordering::SeqCst);
for op in step {
self.execute_op(op)?;
}
}
// Collect outputs from last stage (simulated)
let num_outputs = self.config.num_micro_batches;
for i in 0..num_outputs {
outputs.push(MicroBatch::new(i as u64, i, vec![], vec![]));
}
// Update stats
{
let mut stats = self.stats.write();
stats.mini_batches_processed += 1;
stats.total_time += start.elapsed();
stats.micro_batches_processed += self.config.num_micro_batches;
}
self.running.store(false, Ordering::SeqCst);
self.current_step.store(0, Ordering::SeqCst);
// Reset stages for next iteration
for stage in &self.stages {
stage.reset();
}
Ok(outputs)
}
/// Execute a single pipeline operation
fn execute_op(&self, op: &PipelineOp) -> Result<()> {
match op {
PipelineOp::Forward {
stage,
micro_batch_id,
} => {
let stage_ref = &self.stages[*stage];
stage_ref.set_state(StageState::Forward);
let start = Instant::now();
// Get input tensor:
// - First stage gets from input queue
// - Other stages get from previous stage's output stash
let input_tensor = if *stage == 0 {
// Get from forward queue
if let Some(batch) = stage_ref.dequeue_forward() {
batch.get_tensor()?
} else {
// No input - create dummy tensor for simulation
use rtx_tensor::Device;
Tensor::zeros([1, 64], &Device::cpu()).map_err(|e| {
DistributedError::runtime(format!("Tensor creation failed: {:?}", e))
})?
}
} else {
// Get from previous stage's output
let prev_stage = &self.stages[*stage - 1];
prev_stage
.get_stashed_tensor(*micro_batch_id)
.unwrap_or_else(|| {
// Fallback for simulation
use rtx_tensor::Device;
Tensor::zeros([1, 64], &Device::cpu()).unwrap()
})
};
// Execute actual forward pass using stage module
let output_tensor = stage_ref.execute_forward(&input_tensor)?;
// Stash input activation for backward pass
if self.config.checkpoint_activations {
stage_ref.stash_tensor(*micro_batch_id, input_tensor);
}
// Stash output for next stage (or final output)
if *stage < self.stages.len() - 1 {
// Pass to next stage
self.stages[*stage + 1].stash_tensor(*micro_batch_id, output_tensor.clone());
} else {
// Last stage - stash output for later retrieval
stage_ref.stash_tensor(*micro_batch_id + 1000, output_tensor); // Use offset to distinguish
}
stage_ref.record_forward(start.elapsed());
stage_ref.set_state(StageState::Idle);
}
PipelineOp::Backward {
stage,
micro_batch_id,
} => {
let stage_ref = &self.stages[*stage];
stage_ref.set_state(StageState::Backward);
let start = Instant::now();
// Get gradient from next stage (or loss for last stage)
let grad_output = if *stage == self.stages.len() - 1 {
// Last stage - compute loss gradient (ones for simplicity)
use rtx_tensor::Device;
Tensor::ones([1, 64], &Device::cpu()).map_err(|e| {
DistributedError::runtime(format!("Gradient creation failed: {:?}", e))
})?
} else {
// Get gradient from next stage
let next_stage = &self.stages[*stage + 1];
next_stage
.get_stashed_tensor(*micro_batch_id + 2000) // Gradient offset
.unwrap_or_else(|| {
use rtx_tensor::Device;
Tensor::ones([1, 64], &Device::cpu()).unwrap()
})
};
// Retrieve stashed input activation
let saved_input = stage_ref.get_stashed_tensor(*micro_batch_id);
// Execute actual backward pass using stage module
let grad_input = stage_ref.execute_backward(&grad_output, saved_input.as_ref())?;
// Stash gradient for previous stage
if *stage > 0 {
self.stages[*stage - 1].stash_tensor(*micro_batch_id + 2000, grad_input); // Gradient offset
}
// Clean up stashed activation
stage_ref.remove_stashed_tensor(*micro_batch_id);
stage_ref.record_backward(start.elapsed());
stage_ref.set_state(StageState::Idle);
}
PipelineOp::SendActivation {
from_stage,
to_stage,
micro_batch_id,
} => {
// Transfer tensor between stages (already done via stash in Forward)
// In distributed setting, this would involve actual P2P communication
if !self.config.overlap_comm {
// Simulate transfer time if not overlapped
std::thread::sleep(Duration::from_micros(10));
}
let _ = (from_stage, to_stage, micro_batch_id);
}
PipelineOp::RecvActivation {
from_stage,
to_stage,
micro_batch_id,
} => {
// Tensor already available via stash
let _ = (from_stage, to_stage, micro_batch_id);
}
PipelineOp::SendGradient {
from_stage,
to_stage,
micro_batch_id,
} => {
// Transfer gradient between stages (already done via stash in Backward)
if !self.config.overlap_comm {
std::thread::sleep(Duration::from_micros(10));
}
let _ = (from_stage, to_stage, micro_batch_id);
}
PipelineOp::RecvGradient {
from_stage,
to_stage,
micro_batch_id,
} => {
// Gradient already available via stash
let _ = (from_stage, to_stage, micro_batch_id);
}
PipelineOp::Barrier => {
// Synchronization point - in distributed setting would use collective
}
}
Ok(())
}
/// Set module for a specific stage
pub fn set_stage_module(&self, stage_id: usize, module: Arc<dyn StageModule>) -> Result<()> {
if stage_id >= self.stages.len() {
return Err(DistributedError::configuration(format!(
"Stage {} does not exist (max: {})",
stage_id,
self.stages.len() - 1
)));
}
self.stages[stage_id].set_module(module);
Ok(())
}
/// Zero all gradients across stages
pub fn zero_grad(&self) {
for stage in &self.stages {
stage.zero_grad();
}
}
/// Apply gradients with learning rate
pub fn apply_gradients(&self, lr: f32) {
for stage in &self.stages {
stage.apply_gradients(lr);
}
}
/// Get stage by ID
pub fn stage(&self, stage_id: usize) -> Option<Arc<PipelineStage>> {
self.stages.get(stage_id).cloned()
}
/// Get all stage statistics
pub fn stage_stats(&self) -> Vec<StageStats> {
self.stages.iter().map(|s| s.stats()).collect()
}
/// Get pipeline statistics
pub fn stats(&self) -> PipelineStats {
self.stats.read().clone()
}
/// Check if pipeline is running
pub fn is_running(&self) -> bool {
self.running.load(Ordering::SeqCst)
}
/// Get current step
pub fn current_step(&self) -> usize {
self.current_step.load(Ordering::SeqCst)
}
/// Get number of stages
pub fn num_stages(&self) -> usize {
self.stages.len()
}
}
/// Thread-safe shared pipeline executor
pub type SharedPipelineExecutor = Arc<PipelineExecutor>;
/// Create a shared pipeline executor
pub fn shared_pipeline_executor(
config: PipelineConfig,
device_ids: Vec<i32>,
) -> Result<SharedPipelineExecutor> {
Ok(Arc::new(PipelineExecutor::new(config, device_ids)?))
}