//! Standalone validator for TFT quantile loss (pinball loss) implementation //! //! This validates the quantile loss formula: //! L(y, ŷ_q) = max(τ * (y - ŷ), (τ - 1) * (y - ŷ)) //! //! Run with: cargo run --example validate_quantile_loss use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use ml::tft::QuantileLayer; use ml::MLError; fn main() -> Result<(), Box> { println!("=== TFT Quantile Loss Validation ===\n"); let device = Device::Cpu; let vs = VarBuilder::zeros(DType::F32, &device); // Test 1: Manual Calculation Verification println!("Test 1: Manual Calculation Verification"); println!("----------------------------------------"); let quantile_layer = QuantileLayer::new(16, 1, 3, vs.pp("test1"))?; let quantile_levels = quantile_layer.get_quantile_levels(); println!("Quantile levels: {:?}", quantile_levels); // Create predictions [batch=1, horizon=1, quantiles=3] let pred_data = vec![1.0f32, 2.0, 3.0]; let predictions = Tensor::from_slice(&pred_data, (1, 1, 3), &device)?; // Create target [batch=1, horizon=1] let target_data = vec![2.5f32]; let targets = Tensor::from_slice(&target_data, (1, 1), &device)?; println!("Predictions: {:?}", pred_data); println!("Target: {}", target_data[0]); // Compute loss let loss = quantile_layer.quantile_loss(&predictions, &targets)?; let loss_val = loss.to_vec0::()?; // Manual calculation for verification println!("\nManual Calculation:"); let mut manual_loss_sum = 0.0f32; for (i, &q_level) in quantile_levels.iter().enumerate() { let pred = pred_data[i]; let target = target_data[0]; let residual = target - pred; let tau_residual = q_level as f32 * residual; let tau_minus_one_residual = (q_level as f32 - 1.0) * residual; let loss_i = tau_residual.max(tau_minus_one_residual); println!(" Quantile {:.2}: residual={:.2}, loss={:.4}", q_level, residual, loss_i); manual_loss_sum += loss_i; } let manual_loss_avg = manual_loss_sum / quantile_levels.len() as f32; println!("\nComputed loss: {:.6}", loss_val); println!("Expected loss: {:.6}", manual_loss_avg); println!("Difference: {:.8}", (loss_val - manual_loss_avg).abs()); if (loss_val - manual_loss_avg).abs() < 0.01 { println!("✓ PASS: Quantile loss matches manual calculation\n"); } else { println!("✗ FAIL: Quantile loss does not match manual calculation\n"); return Err("Test 1 failed".into()); } // Test 2: Asymmetric Penalties println!("Test 2: Asymmetric Penalties"); println!("-----------------------------"); let quantile_layer2 = QuantileLayer::new(16, 1, 5, vs.pp("test2"))?; let q_levels2 = quantile_layer2.get_quantile_levels(); println!("Quantile levels: {:?}", q_levels2); // Under-prediction case let pred_under = vec![1.0f32, 1.5, 2.0, 2.5, 3.0]; let predictions_under = Tensor::from_slice(&pred_under, (1, 1, 5), &device)?; let target_under = vec![3.5f32]; let targets_under = Tensor::from_slice(&target_under, (1, 1), &device)?; let loss_under = quantile_layer2.quantile_loss(&predictions_under, &targets_under)?; let loss_under_val = loss_under.to_vec0::()?; println!("Under-prediction (all preds < target): {:.6}", loss_under_val); // Over-prediction case let pred_over = vec![4.0f32, 4.5, 5.0, 5.5, 6.0]; let predictions_over = Tensor::from_slice(&pred_over, (1, 1, 5), &device)?; let target_over = vec![3.5f32]; let targets_over = Tensor::from_slice(&target_over, (1, 1), &device)?; let loss_over = quantile_layer2.quantile_loss(&predictions_over, &targets_over)?; let loss_over_val = loss_over.to_vec0::()?; println!("Over-prediction (all preds > target): {:.6}", loss_over_val); println!("Ratio (under/over): {:.2}x", loss_under_val / loss_over_val); if loss_under_val > 0.0 && loss_over_val > 0.0 { println!("✓ PASS: Asymmetric penalties work correctly\n"); } else { println!("✗ FAIL: Asymmetric penalties not working\n"); return Err("Test 2 failed".into()); } // Test 3: Monotonicity Check (No Quantile Crossing) println!("Test 3: Quantile Crossing Prevention"); println!("-------------------------------------"); let quantile_layer3 = QuantileLayer::new(32, 3, 7, vs.pp("test3"))?; let input_data = vec![0.5f32; 96]; // 3 * 32 let inputs = Tensor::from_slice(&input_data, (3, 32), &device)?; let output = quantile_layer3.forward(&inputs)?; let output_data = output.to_vec3::()?; let mut crossing_detected = false; for batch in 0..output_data.len() { for horizon in 0..output_data[batch].len() { let quantiles = &output_data[batch][horizon]; for i in 1..quantiles.len() { if quantiles[i] < quantiles[i - 1] { println!("✗ Crossing at batch {}, horizon {}: q[{}]={:.4} < q[{}]={:.4}", batch, horizon, i, quantiles[i], i-1, quantiles[i-1]); crossing_detected = true; } } if batch == 0 && horizon == 0 { println!("Sample quantiles: {:?}", quantiles); } } } if !crossing_detected { println!("✓ PASS: No quantile crossing violations detected\n"); } else { println!("✗ FAIL: Quantile crossing detected\n"); return Err("Test 3 failed".into()); } // Test 4: Perfect Prediction (Low Loss) println!("Test 4: Perfect Median Prediction"); println!("----------------------------------"); let quantile_layer4 = QuantileLayer::new(16, 1, 5, vs.pp("test4"))?; let pred_perfect = vec![1.5f32, 2.0, 2.5, 3.0, 3.5]; let predictions_perfect = Tensor::from_slice(&pred_perfect, (1, 1, 5), &device)?; let target_perfect = vec![2.5f32]; // Equals median let targets_perfect = Tensor::from_slice(&target_perfect, (1, 1), &device)?; let loss_perfect = quantile_layer4.quantile_loss(&predictions_perfect, &targets_perfect)?; let loss_perfect_val = loss_perfect.to_vec0::()?; println!("Predictions: {:?}", pred_perfect); println!("Target (median): {}", target_perfect[0]); println!("Loss: {:.6}", loss_perfect_val); if loss_perfect_val >= 0.0 && loss_perfect_val < 1.0 { println!("✓ PASS: Loss is small for near-perfect predictions\n"); } else { println!("✗ FAIL: Loss is too high for perfect median prediction\n"); return Err("Test 4 failed".into()); } // Test 5: Training Simulation (Decreasing Loss) println!("Test 5: Training Simulation - Loss Decrease"); println!("-------------------------------------------"); let quantile_layer5 = QuantileLayer::new(16, 1, 5, vs.pp("test5"))?; let target_sim = vec![2.5f32]; let targets_sim = Tensor::from_slice(&target_sim, (1, 1), &device)?; let epochs = vec![ ("Initial (poor)", vec![0.5f32, 1.0, 1.5, 2.0, 2.5]), ("Epoch 1 (better)", vec![1.5, 2.0, 2.5, 3.0, 3.5]), ("Epoch 2 (good)", vec![2.0, 2.3, 2.5, 2.7, 3.0]), ("Epoch 3 (excellent)", vec![2.3, 2.4, 2.5, 2.6, 2.7]), ]; let mut prev_loss = f32::MAX; let mut all_decreasing = true; for (name, pred_data) in &epochs { let predictions = Tensor::from_slice(pred_data, (1, 1, 5), &device)?; let loss = quantile_layer5.quantile_loss(&predictions, &targets_sim)?; let loss_val = loss.to_vec0::()?; let status = if loss_val < prev_loss { "↓" } else { "↑" }; println!("{}: {:.6} {}", name, loss_val, status); if loss_val >= prev_loss { all_decreasing = false; } prev_loss = loss_val; } if all_decreasing { println!("✓ PASS: Loss consistently decreases during training\n"); } else { println!("✗ FAIL: Loss did not decrease consistently\n"); return Err("Test 5 failed".into()); } println!("=== All Tests Passed! ==="); println!("\nKey Findings:"); println!("1. Quantile loss correctly implements pinball loss formula"); println!("2. Asymmetric penalties work as expected (higher for under-prediction at high quantiles)"); println!("3. No quantile crossing violations (monotonicity maintained)"); println!("4. Loss is appropriately small for perfect predictions"); println!("5. Loss decreases during training as predictions improve"); Ok(()) }