Files
rustytorch/demos/rtx-mre/src/training.rs
T
2026-03-04 00:08:42 +00:00

173 lines
4.7 KiB
Rust

//! Training utilities for MRE solver
//!
//! Provides high-level training functions and progress tracking.
use crate::config::MreConfig;
use crate::phantom::PhantomGenerator;
use crate::solver::MreSolver;
use anyhow::Result;
use mre_shared::{LossRecord, MreSnapshot, PhantomConfig};
/// Training session configuration
#[derive(Debug, Clone)]
pub struct TrainingConfig {
/// Number of training steps
pub num_steps: usize,
/// Snapshot interval (every N steps)
pub snapshot_interval: usize,
/// Early stopping threshold for loss
pub early_stop_threshold: Option<f32>,
/// Print progress every N steps
pub print_interval: usize,
}
impl Default for TrainingConfig {
fn default() -> Self {
Self {
num_steps: 1000,
snapshot_interval: 100,
early_stop_threshold: Some(1e-6),
print_interval: 100,
}
}
}
/// Training session result
#[derive(Debug)]
pub struct TrainingResult {
/// All loss records
pub losses: Vec<LossRecord>,
/// Snapshots taken during training
pub snapshots: Vec<MreSnapshot>,
/// Final snapshot
pub final_snapshot: MreSnapshot,
/// Whether training converged
pub converged: bool,
/// Total training time in seconds
pub training_time_secs: f32,
}
/// Run a complete training session
pub fn train_solver(
solver: &mut MreSolver,
training_config: &TrainingConfig,
) -> Result<TrainingResult> {
let start_time = std::time::Instant::now();
let mut losses = Vec::with_capacity(training_config.num_steps);
let mut snapshots = Vec::new();
let mut converged = false;
let snapshot_res = solver.config().stiffness_nx.min(64);
for i in 0..training_config.num_steps {
// Training step
let loss = solver.step()?;
losses.push(loss.clone());
// Check for early stopping
if let Some(threshold) = training_config.early_stop_threshold
&& loss.total_loss < threshold
{
converged = true;
break;
}
// Take snapshot
if (i + 1) % training_config.snapshot_interval == 0 {
let snapshot = solver.snapshot(snapshot_res, snapshot_res)?;
snapshots.push(snapshot);
}
// Print progress
if (i + 1) % training_config.print_interval == 0 {
println!(
"Step {}: total={:.6}, physics={:.6}, data={:.6}",
i + 1,
loss.total_loss,
loss.physics_loss,
loss.data_loss
);
}
}
let final_snapshot = solver.snapshot(snapshot_res, snapshot_res)?;
let training_time = start_time.elapsed().as_secs_f32();
Ok(TrainingResult {
losses,
snapshots,
final_snapshot,
converged,
training_time_secs: training_time,
})
}
/// Quick validation test with synthetic phantom
pub fn validate_with_phantom(mre_config: &MreConfig) -> Result<TrainingResult> {
let mut solver = MreSolver::new(mre_config.clone())?;
// Generate phantom
let phantom = PhantomGenerator::tumor_phantom(mre_config.clone());
let (_ground_truth, wave) = phantom.generate();
solver.set_measured_wave(wave);
// Train
let training_config = TrainingConfig {
num_steps: 500,
snapshot_interval: 100,
early_stop_threshold: Some(1e-5),
print_interval: 50,
};
train_solver(&mut solver, &training_config)
}
/// Create solver from phantom configuration
pub fn create_solver_with_phantom(
mre_config: MreConfig,
phantom_config: PhantomConfig,
) -> Result<MreSolver> {
let mut solver = MreSolver::new(mre_config.clone())?;
let phantom = PhantomGenerator::new(phantom_config, mre_config);
let (_stiffness, wave) = phantom.generate();
solver.set_measured_wave(wave);
Ok(solver)
}
#[cfg(test)]
mod tests {
use super::*;
fn get_test_config() -> MreConfig {
MreConfig::fast()
.with_stiffness_resolution(8, 8)
.with_wave_net_layers(2)
.with_wave_net_hidden(16)
.with_fourier_features(4)
}
#[test]
fn test_training_session() {
let mre_config = get_test_config();
let phantom_config = PhantomConfig::single_tumor();
let mut solver = create_solver_with_phantom(mre_config, phantom_config).unwrap();
let training_config = TrainingConfig {
num_steps: 10,
snapshot_interval: 5,
early_stop_threshold: None,
print_interval: 5,
};
let result = train_solver(&mut solver, &training_config).unwrap();
assert_eq!(result.losses.len(), 10);
assert!(result.snapshots.len() >= 1);
}
}