Initial commit
This commit is contained in:
@@ -0,0 +1,491 @@
|
||||
//! 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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user