Files
rustytorch/crates/training/rtx-nas/src/hardware/latency.rs
T
2026-03-04 00:08:42 +00:00

492 lines
15 KiB
Rust

//! Latency prediction for hardware-aware NAS
//!
//! Provides latency estimation and prediction for architectures
//! on different hardware devices.
use super::device::DeviceProfile;
use crate::error::{NASError, Result};
use crate::search_space::{Architecture, Cell, OperationType};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// Latency predictor trait
pub trait LatencyPredictor: Send + Sync {
/// Predict latency for an architecture on a device
fn predict(&self, arch: &Architecture, device: &DeviceProfile) -> Result<f32>;
/// Calibrate the predictor with measured samples
fn calibrate(&mut self, samples: &[(Architecture, f32)]) -> Result<()>;
/// Get prediction confidence (0.0 to 1.0)
fn confidence(&self) -> f32;
}
/// Operation-level latency lookup table
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OperationLatency {
/// Base latency in microseconds
pub base_latency_us: f32,
/// Latency per channel
pub per_channel_us: f32,
/// Latency per spatial element (H*W)
pub per_spatial_us: f32,
/// Memory access pattern factor
pub memory_factor: f32,
}
impl OperationLatency {
/// Create a new operation latency entry
pub fn new(base: f32, per_channel: f32, per_spatial: f32) -> Self {
Self {
base_latency_us: base,
per_channel_us: per_channel,
per_spatial_us: per_spatial,
memory_factor: 1.0,
}
}
/// Estimate latency for given configuration
pub fn estimate(&self, channels: usize, height: usize, width: usize) -> f32 {
let spatial = (height * width) as f32;
self.base_latency_us
+ self.per_channel_us * channels as f32
+ self.per_spatial_us * spatial * self.memory_factor
}
}
/// Lookup table based latency predictor
#[derive(Debug, Clone)]
pub struct LookupTablePredictor {
/// Latency tables per device type per operation
tables: HashMap<String, HashMap<OperationType, OperationLatency>>,
/// Default latencies when no specific table exists
default_latencies: HashMap<OperationType, OperationLatency>,
/// Confidence in predictions
confidence: f32,
}
impl LookupTablePredictor {
/// Create a new lookup table predictor with default values
pub fn new() -> Self {
let mut default_latencies = HashMap::new();
// Default latencies based on typical GPU performance
default_latencies.insert(
OperationType::Identity,
OperationLatency::new(0.1, 0.0, 0.0001),
);
default_latencies.insert(OperationType::Zero, OperationLatency::new(0.05, 0.0, 0.0));
default_latencies.insert(
OperationType::Conv3x3,
OperationLatency::new(10.0, 0.5, 0.01),
);
default_latencies.insert(
OperationType::Conv5x5,
OperationLatency::new(25.0, 1.2, 0.025),
);
default_latencies.insert(
OperationType::SepConv3x3,
OperationLatency::new(8.0, 0.3, 0.008),
);
default_latencies.insert(
OperationType::SepConv5x5,
OperationLatency::new(15.0, 0.6, 0.015),
);
default_latencies.insert(
OperationType::DilConv3x3,
OperationLatency::new(12.0, 0.6, 0.012),
);
default_latencies.insert(
OperationType::MaxPool3x3,
OperationLatency::new(2.0, 0.1, 0.002),
);
default_latencies.insert(
OperationType::AvgPool3x3,
OperationLatency::new(2.5, 0.1, 0.0025),
);
Self {
tables: HashMap::new(),
default_latencies,
confidence: 0.5, // Medium confidence for uncalibrated predictor
}
}
/// Add device-specific latency table
pub fn add_device_table(
&mut self,
device_name: impl Into<String>,
table: HashMap<OperationType, OperationLatency>,
) {
self.tables.insert(device_name.into(), table);
}
/// Get latency for an operation on a device
fn get_op_latency(&self, op: OperationType, device: &DeviceProfile) -> OperationLatency {
// Try device-specific table first
if let Some(device_table) = self.tables.get(&device.name) {
if let Some(latency) = device_table.get(&op) {
return latency.clone();
}
}
// Fall back to default
self.default_latencies
.get(&op)
.cloned()
.unwrap_or_else(|| OperationLatency::new(10.0, 0.5, 0.01))
}
/// Estimate latency for a cell
fn estimate_cell_latency(
&self,
cell: &Cell,
device: &DeviceProfile,
channels: usize,
height: usize,
width: usize,
) -> f32 {
let mut total_latency = 0.0;
for edge in cell.edges() {
if let Some(op_type) = cell.get_operation(&edge) {
let op_latency = self.get_op_latency(op_type, device);
total_latency += op_latency.estimate(channels, height, width);
}
}
// Apply device-specific scaling
let device_factor = match device.device_type {
super::device::DeviceType::CPU => 10.0,
super::device::DeviceType::Mobile => 5.0,
super::device::DeviceType::Edge => 20.0,
_ => 1.0, // GPUs as baseline
};
total_latency * device_factor / device.peak_tflops_fp32.max(0.1)
}
}
impl Default for LookupTablePredictor {
fn default() -> Self {
Self::new()
}
}
impl LatencyPredictor for LookupTablePredictor {
fn predict(&self, arch: &Architecture, device: &DeviceProfile) -> Result<f32> {
let mut total_latency_ms = 0.0;
// Default dimensions if not specified in architecture
let channels = arch.channels;
let height = 224; // Default ImageNet size
let width = 224;
for cell in &arch.cells {
let cell_latency_us = self.estimate_cell_latency(cell, device, channels, height, width);
total_latency_ms += cell_latency_us / 1000.0; // Convert to ms
}
Ok(total_latency_ms)
}
fn calibrate(&mut self, samples: &[(Architecture, f32)]) -> Result<()> {
if samples.is_empty() {
return Err(NASError::LatencyError(
"No calibration samples provided".into(),
));
}
// Simple calibration: adjust confidence based on sample count
let sample_count = samples.len();
self.confidence = (sample_count as f32 / 100.0).min(0.95);
// In a real implementation, we would fit the latency tables to the samples
Ok(())
}
fn confidence(&self) -> f32 {
self.confidence
}
}
/// Linear regression based latency predictor
#[derive(Debug, Clone)]
pub struct RegressionPredictor {
/// Learned coefficients for each operation type
coefficients: HashMap<OperationType, Vec<f32>>,
/// Bias term
bias: f32,
/// Prediction confidence
confidence: f32,
/// Number of features per operation
num_features: usize,
}
impl RegressionPredictor {
/// Create a new regression predictor
pub fn new() -> Self {
Self {
coefficients: HashMap::new(),
bias: 0.0,
confidence: 0.0, // No confidence until trained
num_features: 4, // channels, height, width, device_factor
}
}
/// Extract features from an architecture
fn extract_features(&self, arch: &Architecture, device: &DeviceProfile) -> Vec<f32> {
let mut features = Vec::new();
// Count operations of each type
let mut op_counts: HashMap<OperationType, usize> = HashMap::new();
for cell in &arch.cells {
for edge in cell.edges() {
if let Some(op_type) = cell.get_operation(&edge) {
*op_counts.entry(op_type).or_insert(0) += 1;
}
}
}
// Add operation counts as features
for op_type in OperationType::all() {
let count = *op_counts.get(&op_type).unwrap_or(&0);
features.push(count as f32);
}
// Add architecture properties
features.push(arch.num_cells() as f32);
features.push(arch.channels as f32);
features.push(arch.channels as f32); // Same for in/out
// Add device factor
let device_factor = 1.0 / device.peak_tflops_fp32.max(0.1);
features.push(device_factor);
features
}
/// Predict using learned coefficients
fn predict_internal(&self, features: &[f32]) -> f32 {
if self.coefficients.is_empty() {
// Fall back to simple estimation
return features.iter().sum::<f32>() * 0.1 + self.bias;
}
// Linear combination of features
let mut prediction = self.bias;
let num_ops = OperationType::count();
for (i, op_type) in OperationType::all().iter().enumerate() {
if let Some(coeffs) = self.coefficients.get(op_type) {
if i < features.len() {
prediction += features[i] * coeffs.first().unwrap_or(&1.0);
}
}
}
// Add contribution from other features
for (_i, &f) in features.iter().enumerate().skip(num_ops) {
prediction += f * 0.1; // Simple scaling for non-operation features
}
prediction.max(0.0) // Latency can't be negative
}
}
impl Default for RegressionPredictor {
fn default() -> Self {
Self::new()
}
}
impl LatencyPredictor for RegressionPredictor {
fn predict(&self, arch: &Architecture, device: &DeviceProfile) -> Result<f32> {
let features = self.extract_features(arch, device);
Ok(self.predict_internal(&features))
}
fn calibrate(&mut self, samples: &[(Architecture, f32)]) -> Result<()> {
if samples.is_empty() {
return Err(NASError::LatencyError(
"No calibration samples provided".into(),
));
}
// Simple least squares fitting
// In a real implementation, use proper linear regression
// Compute mean latency as baseline
let mean_latency: f32 = samples.iter().map(|(_, l)| l).sum::<f32>() / samples.len() as f32;
self.bias = mean_latency;
// Initialize coefficients for each operation
for op_type in OperationType::all() {
self.coefficients
.insert(op_type, vec![1.0; self.num_features]);
}
// Update confidence based on sample count
self.confidence = (samples.len() as f32 / 100.0).min(0.9);
Ok(())
}
fn confidence(&self) -> f32 {
self.confidence
}
}
/// Latency measurement result
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LatencyMeasurement {
/// Architecture identifier
pub arch_id: String,
/// Device name
pub device_name: String,
/// Measured latency in milliseconds
pub latency_ms: f32,
/// Standard deviation (if multiple measurements)
pub std_dev: Option<f32>,
/// Number of measurements
pub num_samples: usize,
/// Batch size used for measurement
pub batch_size: usize,
}
impl LatencyMeasurement {
/// Create a new latency measurement
pub fn new(
arch_id: impl Into<String>,
device_name: impl Into<String>,
latency_ms: f32,
) -> Self {
Self {
arch_id: arch_id.into(),
device_name: device_name.into(),
latency_ms,
std_dev: None,
num_samples: 1,
batch_size: 1,
}
}
/// Add statistical information
pub fn with_stats(mut self, std_dev: f32, num_samples: usize) -> Self {
self.std_dev = Some(std_dev);
self.num_samples = num_samples;
self
}
/// Set batch size
pub fn with_batch_size(mut self, batch_size: usize) -> Self {
self.batch_size = batch_size;
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::search_space::CellConfig;
fn create_test_architecture() -> Architecture {
let config = CellConfig::default_darts();
let cell = Cell::new(config).unwrap();
Architecture::new("test_arch".to_string(), vec![cell], 16, 32)
}
#[test]
fn test_operation_latency() {
let latency = OperationLatency::new(10.0, 0.5, 0.01);
let estimate = latency.estimate(64, 32, 32);
assert!(estimate > 10.0); // Should be more than base
}
#[test]
fn test_lookup_table_predictor_new() {
let predictor = LookupTablePredictor::new();
assert_eq!(predictor.confidence(), 0.5);
assert!(!predictor.default_latencies.is_empty());
}
#[test]
fn test_lookup_table_predict() {
let predictor = LookupTablePredictor::new();
let arch = create_test_architecture();
let device = super::super::device::CommonDevices::rtx_3090();
let result = predictor.predict(&arch, &device);
assert!(result.is_ok());
assert!(result.unwrap() >= 0.0);
}
#[test]
fn test_lookup_table_calibrate() {
let mut predictor = LookupTablePredictor::new();
let arch = create_test_architecture();
let samples = vec![
(arch.clone(), 1.0),
(arch.clone(), 1.1),
(arch.clone(), 0.9),
];
let result = predictor.calibrate(&samples);
assert!(result.is_ok());
assert!(predictor.confidence() > 0.0);
}
#[test]
fn test_regression_predictor_new() {
let predictor = RegressionPredictor::new();
assert_eq!(predictor.confidence(), 0.0);
}
#[test]
fn test_regression_predict() {
let predictor = RegressionPredictor::new();
let arch = create_test_architecture();
let device = super::super::device::CommonDevices::rtx_3090();
let result = predictor.predict(&arch, &device);
assert!(result.is_ok());
assert!(result.unwrap() >= 0.0);
}
#[test]
fn test_regression_calibrate() {
let mut predictor = RegressionPredictor::new();
let arch = create_test_architecture();
let samples = vec![
(arch.clone(), 1.0),
(arch.clone(), 1.1),
(arch.clone(), 0.9),
];
let result = predictor.calibrate(&samples);
assert!(result.is_ok());
assert!(predictor.confidence() > 0.0);
}
#[test]
fn test_latency_measurement() {
let measurement = LatencyMeasurement::new("arch_1", "RTX 3090", 1.5)
.with_stats(0.1, 10)
.with_batch_size(1);
assert_eq!(measurement.arch_id, "arch_1");
assert_eq!(measurement.latency_ms, 1.5);
assert_eq!(measurement.std_dev, Some(0.1));
assert_eq!(measurement.num_samples, 10);
}
#[test]
fn test_calibrate_empty_samples() {
let mut predictor = LookupTablePredictor::new();
let result = predictor.calibrate(&[]);
assert!(result.is_err());
}
}