# MLTrainingService API Documentation ## Overview The MLTrainingService is a core component of the Foxhunt HFT system that provides comprehensive machine learning model training capabilities with enterprise-grade safety controls, real-time monitoring, and production-ready workflow management. ## Table of Contents 1. [Service Architecture](#service-architecture) 2. [gRPC Service Definition](#grpc-service-definition) 3. [Training Workflow](#training-workflow) 4. [API Reference](#api-reference) 5. [Data Pipeline](#data-pipeline) 6. [Lifecycle Management](#lifecycle-management) 7. [Monitoring & Observability](#monitoring--observability) 8. [Integration Examples](#integration-examples) 9. [Error Handling](#error-handling) 10. [Production Deployment](#production-deployment) ## Service Architecture The MLTrainingService follows a microservice architecture pattern with the following components: ``` ┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐ │ TLI Client │ │ MLTraining │ │ Training │ │ (gRPC) │◄──►│ Service │◄──►│ Pipeline │ └─────────────────┘ └─────────────────┘ └─────────────────┘ │ │ ▼ ▼ ┌─────────────────┐ ┌─────────────────┐ │ Model Registry │ │ Safety Manager │ │ (DashMap) │ │ (Gradient/NaN) │ └─────────────────┘ └─────────────────┘ │ │ ▼ ▼ ┌─────────────────┐ ┌─────────────────┐ │ Persistence │ │ Resource Mgmt │ │ (PostgreSQL) │ │ (GPU/CPU) │ └─────────────────┘ └─────────────────┘ ``` ### Key Features - **Real-time Training Monitoring**: Streaming progress updates with sub-second latency - **Enterprise Safety Controls**: Mathematical safety guarantees, gradient clipping, NaN detection - **Multi-Model Support**: DQN, PPO, MAMBA, TFT, Liquid Neural Networks, Transformers - **Resource Management**: Dynamic GPU/CPU allocation with utilization monitoring - **Financial Type Safety**: Unified decimal types preventing precision loss - **Automatic Deployment**: Optional auto-deployment upon successful training completion ## gRPC Service Definition The MLTrainingService is defined in `/tli/proto/ml.proto` and provides the following service interface: ```protobuf service MLTrainingService { // Training job management rpc StartTraining(StartTrainingRequest) returns (TrainingJob); rpc StopTraining(StopTrainingRequest) returns (TrainingJob); rpc ListTrainingJobs(ListTrainingJobsRequest) returns (ListTrainingJobsResponse); // Real-time training monitoring (streaming) rpc WatchTrainingProgress(WatchTrainingRequest) returns (stream TrainingProgressUpdate); // Training configuration and validation rpc ValidateTrainingConfig(TrainingConfigRequest) returns (TrainingConfigResponse); rpc GetTrainingTemplates(TrainingTemplatesRequest) returns (TrainingTemplatesResponse); // Resource management rpc GetResourceUtilization(ResourceRequest) returns (ResourceResponse); rpc StreamResourceMetrics(ResourceRequest) returns (stream ResourceMetricsUpdate); } ``` ## Training Workflow ### 1. Training Job Lifecycle ```mermaid graph TD A[Submit Training Request] --> B[Validate Configuration] B --> C{Configuration Valid?} C -->|No| D[Return Validation Errors] C -->|Yes| E[Allocate Resources] E --> F[Load Dataset] F --> G[Initialize Model] G --> H[Start Training Loop] H --> I[Monitor Progress] I --> J{Training Complete?} J -->|No| K[Update Metrics] K --> H J -->|Yes| L[Validate Model] L --> M{Auto Deploy?} M -->|Yes| N[Deploy Model] M -->|No| O[Store Model] N --> P[Cleanup Resources] O --> P P --> Q[Training Complete] ``` ### 2. Training States | State | Description | Next Possible States | |-------|-------------|---------------------| | `QUEUED` | Training job submitted and waiting for resources | `PREPARING`, `CANCELLED` | | `PREPARING` | Allocating resources and loading data | `RUNNING`, `FAILED` | | `RUNNING` | Active training in progress | `COMPLETED`, `FAILED`, `STOPPING` | | `STOPPING` | Graceful shutdown in progress | `CANCELLED`, `FAILED` | | `COMPLETED` | Training finished successfully | Terminal state | | `FAILED` | Training failed with error | Terminal state | | `CANCELLED` | Training was cancelled by user | Terminal state | ## API Reference ### StartTraining Initiates a new training job with comprehensive validation and resource allocation. **Request:** ```protobuf message StartTrainingRequest { string model_name = 1; // "DQN_EURUSD_v2" string dataset_id = 2; // "market_data_2024_q1" TrainingHyperparameters hyperparameters = 3; // Training configuration ResourceRequirements resource_requirements = 4; // GPU/CPU requirements repeated string tags = 5; // ["production", "eurusd"] string description = 6; // Human description bool auto_deploy = 7; // Auto-deploy on success } ``` **Response:** ```protobuf message TrainingJob { string job_id = 1; // "train_550e8400-e29b-41d4-a716-446655440000" string model_name = 2; // Echo from request TrainingStatus status = 3; // Current status int64 start_time = 4; // Unix timestamp nanoseconds // ... additional fields } ``` **Example Usage:** ```rust use tli::ml_training_service_client::MlTrainingServiceClient; let mut client = MlTrainingServiceClient::connect("http://localhost:50051").await?; let request = tonic::Request::new(StartTrainingRequest { model_name: "DQN_EURUSD_Production".to_string(), dataset_id: "market_data_2024_q3".to_string(), hyperparameters: Some(TrainingHyperparameters { learning_rate: 0.001, batch_size: 128, epochs: 1000, dropout_rate: Some(0.1), ..Default::default() }), resource_requirements: Some(ResourceRequirements { gpu_count: 1, cpu_cores: 8, memory_gb: 16, gpu_type: Some("A100".to_string()), disk_gb: 100, }), tags: vec!["production".to_string(), "eurusd".to_string()], description: "Production DQN training for EURUSD pair".to_string(), auto_deploy: true, }); let response = client.start_training(request).await?; let job = response.into_inner(); println!("Training job started: {}", job.job_id); ``` ### WatchTrainingProgress Streams real-time training progress updates with metrics, logs, and resource utilization. **Request:** ```protobuf message WatchTrainingRequest { string job_id = 1; // Training job to monitor bool include_logs = 2; // Include log messages bool include_metrics = 3; // Include training metrics } ``` **Response Stream:** ```protobuf message TrainingProgressUpdate { string job_id = 1; TrainingStatus status = 2; int32 current_epoch = 3; int32 total_epochs = 4; double progress_percentage = 5; // 0.0 to 100.0 TrainingMetrics metrics = 6; // Loss, accuracy, etc. optional string log_message = 7; // Log output int64 timestamp = 8; optional ResourceUtilization resource_usage = 9; } ``` **Example Usage:** ```rust let request = tonic::Request::new(WatchTrainingRequest { job_id: job.job_id.clone(), include_logs: true, include_metrics: true, }); let mut stream = client.watch_training_progress(request).await?.into_inner(); while let Some(update) = stream.next().await { let update = update?; println!("Epoch {}/{}: {:.2}% complete", update.current_epoch, update.total_epochs, update.progress_percentage ); if let Some(metrics) = update.metrics { println!("Loss: {:.4}, Accuracy: {:.2}%", metrics.loss, metrics.accuracy * 100.0 ); } if let Some(log) = update.log_message { println!("Log: {}", log); } } ``` ### ListTrainingJobs Retrieves training jobs with filtering and pagination support. **Request:** ```protobuf message ListTrainingJobsRequest { optional string model_name = 1; // Filter by model optional TrainingStatus status = 2; // Filter by status optional int64 start_time_after = 3; // Filter by start time optional int64 start_time_before = 4; repeated string tags = 5; // Filter by tags int32 limit = 6; // Max results (default: 50) string cursor = 7; // Pagination cursor } ``` **Response:** ```protobuf message ListTrainingJobsResponse { repeated TrainingJob jobs = 1; string next_cursor = 2; // For pagination int32 total_count = 3; // Total matching jobs } ``` ### ValidateTrainingConfig Validates training configuration before job submission with suggestions for optimization. **Request:** ```protobuf message TrainingConfigRequest { string model_name = 1; TrainingHyperparameters hyperparameters = 2; ResourceRequirements resource_requirements = 3; } ``` **Response:** ```protobuf message TrainingConfigResponse { bool valid = 1; repeated string validation_errors = 2; repeated string validation_warnings = 3; optional TrainingHyperparameters suggested_params = 4; optional ResourceRequirements suggested_resources = 5; double estimated_duration_hours = 6; } ``` ## Data Pipeline ### Feature Engineering Pipeline The MLTrainingService integrates with a sophisticated feature engineering pipeline: ```rust pub struct FinancialFeatures { /// Price features (normalized, safe decimal representation) pub prices: Vec, /// Volume features (safe integers to prevent overflow) pub volumes: Vec, /// Technical indicators (bounded and validated) pub technical_indicators: HashMap, /// Market microstructure features pub microstructure: MicrostructureFeatures, /// Risk metrics (VaR, Expected Shortfall, etc.) pub risk_metrics: RiskFeatures, /// Timestamp for temporal alignment pub timestamp: chrono::DateTime, } ``` ### Data Validation All training data undergoes comprehensive validation: 1. **Financial Type Safety**: All prices use unified `IntegerPrice` type preventing floating-point precision loss 2. **Range Validation**: Technical indicators bounded to expected ranges 3. **Temporal Consistency**: Timestamps validated for proper chronological order 4. **Missing Data Handling**: Configurable strategies for missing value imputation 5. **Outlier Detection**: Statistical outlier detection with configurable thresholds ### Supported Data Sources - **Real-time Market Data**: Direct integration with Polygon.io and broker feeds - **Historical Data**: PostgreSQL and InfluxDB time-series data - **Alternative Data**: Economic indicators, sentiment data, news feeds - **Custom Datasets**: User-provided datasets with validation ## Lifecycle Management ### Model Versioning The service implements comprehensive model versioning: ```rust pub struct ModelVersion { pub version_id: String, // "v1.2.3" pub model_id: String, // "DQN_EURUSD" pub training_job_id: String, // Reference to training job pub created_at: DateTime, pub performance_metrics: PerformanceMetrics, pub hyperparameters: TrainingHyperparameters, pub deployment_status: DeploymentStatus, } ``` ### Deployment Pipeline ```mermaid graph LR A[Training Complete] --> B[Model Validation] B --> C{Auto Deploy?} C -->|Yes| D[Staging Deployment] C -->|No| E[Model Stored] D --> F[Integration Tests] F --> G{Tests Pass?} G -->|Yes| H[Production Deployment] G -->|No| I[Rollback to Previous] H --> J[Health Monitoring] ``` ### Model Registry Integration Models are automatically registered in the global registry upon successful training: ```rust let registry = get_global_registry(); let trained_model = Arc::new(TLOBModelWrapper::new(tlob_model)); registry.register(trained_model).await?; ``` ## Monitoring & Observability ### Training Metrics Real-time metrics tracked during training: - **Loss Functions**: Training and validation loss with convergence analysis - **Accuracy Metrics**: Precision, recall, F1-score, AUC-ROC - **Financial Metrics**: Sharpe ratio, Calmar ratio, maximum drawdown - **Performance Metrics**: Training speed, GPU utilization, memory usage - **Safety Metrics**: Gradient norms, NaN detection, numerical stability ### Resource Monitoring ```rust pub struct ResourceUtilization { pub gpu_utilization: f64, // 0.0 to 1.0 pub gpu_memory_used: f64, // 0.0 to 1.0 pub cpu_utilization: f64, pub memory_used: f64, pub disk_used: f64, pub timestamp: i64, } ``` ### Alerting Automated alerts for: - Training failures or divergence - Resource exhaustion - Safety violations (NaN, gradient explosion) - Performance degradation - Hardware failures ## Integration Examples ### Basic Training Job ```rust use foxhunt_tli::ml_training::{MLTrainingServiceClient, StartTrainingRequest}; async fn train_dqn_model() -> Result<(), Box> { let mut client = MLTrainingServiceClient::connect("http://localhost:50051").await?; let training_request = StartTrainingRequest { model_name: "DQN_EURUSD_v3".to_string(), dataset_id: "market_data_q3_2024".to_string(), hyperparameters: Some(TrainingHyperparameters { learning_rate: 0.0001, batch_size: 64, epochs: 2000, dropout_rate: Some(0.15), hidden_layers: Some(3), hidden_units: Some(256), custom_params: hashmap! { "epsilon_decay".to_string() => "0.995".to_string(), "target_update_frequency".to_string() => "100".to_string(), }, }), resource_requirements: Some(ResourceRequirements { gpu_count: 1, cpu_cores: 8, memory_gb: 32, gpu_type: Some("A100".to_string()), disk_gb: 200, }), tags: vec!["production".to_string(), "dqn".to_string(), "eurusd".to_string()], description: "Production DQN training for EURUSD with enhanced safety controls".to_string(), auto_deploy: true, }; let response = client.start_training(tonic::Request::new(training_request)).await?; let job = response.into_inner(); println!("Training job started: {} (ID: {})", job.model_name, job.job_id); // Monitor training progress let watch_request = WatchTrainingRequest { job_id: job.job_id.clone(), include_logs: true, include_metrics: true, }; let mut stream = client.watch_training_progress( tonic::Request::new(watch_request) ).await?.into_inner(); while let Some(update) = stream.next().await { let update = update?; match update.status() { TrainingStatus::Running => { if let Some(metrics) = &update.metrics { println!("Epoch {}/{}: Loss={:.4}, Acc={:.2}%, GPU={:.1}%", update.current_epoch, update.total_epochs, metrics.loss, metrics.accuracy * 100.0, update.resource_usage.as_ref().map(|r| r.gpu_utilization * 100.0).unwrap_or(0.0) ); } } TrainingStatus::Completed => { println!("Training completed successfully!"); if let Some(model_id) = &job.resulting_model_id { println!("Model deployed with ID: {}", model_id); } break; } TrainingStatus::Failed => { println!("Training failed: {}", update.log_message.unwrap_or_default()); break; } _ => {} } } Ok(()) } ``` ### Ensemble Training ```rust async fn train_ensemble_model() -> Result<(), Box> { let mut client = MLTrainingServiceClient::connect("http://localhost:50051").await?; // Train individual models for ensemble let base_models = vec!["DQN", "PPO", "MAMBA", "TFT"]; let mut training_jobs = Vec::new(); for model_type in base_models { let request = StartTrainingRequest { model_name: format!("{}_EURUSD_ensemble_base", model_type), dataset_id: "market_data_ensemble_2024".to_string(), hyperparameters: Some(get_model_hyperparameters(model_type)), resource_requirements: Some(get_resource_requirements(model_type)), tags: vec!["ensemble".to_string(), "base_model".to_string()], description: format!("Base {} model for ensemble training", model_type), auto_deploy: false, // Don't auto-deploy base models }; let response = client.start_training(tonic::Request::new(request)).await?; training_jobs.push(response.into_inner()); } // Wait for all base models to complete let completed_models = wait_for_training_completion(&mut client, training_jobs).await?; // Train ensemble meta-model let ensemble_request = StartTrainingRequest { model_name: "ENSEMBLE_EURUSD_v1".to_string(), dataset_id: "market_data_ensemble_2024".to_string(), hyperparameters: Some(TrainingHyperparameters { learning_rate: 0.01, batch_size: 32, epochs: 500, custom_params: hashmap! { "base_models".to_string() => completed_models.join(","), "ensemble_method".to_string() => "stacking".to_string(), }, }), resource_requirements: Some(ResourceRequirements { gpu_count: 1, cpu_cores: 16, memory_gb: 64, disk_gb: 500, }), tags: vec!["ensemble".to_string(), "production".to_string()], description: "Ensemble meta-model combining DQN, PPO, MAMBA, and TFT".to_string(), auto_deploy: true, }; let ensemble_job = client.start_training(tonic::Request::new(ensemble_request)).await?; println!("Ensemble training started: {}", ensemble_job.into_inner().job_id); Ok(()) } ``` ## Error Handling ### Error Categories The MLTrainingService defines comprehensive error categories: ```rust pub enum ProductionTrainingError { ConfigError { reason: String }, // Configuration validation errors ArchitectureError { reason: String }, // Model architecture issues DataError { reason: String }, // Data loading/validation errors OptimizationError { reason: String }, // Training optimization failures FinancialError { reason: String }, // Financial type validation errors SafetyViolation { reason: String }, // Safety control violations ConvergenceError { reason: String }, // Model convergence failures ResourceError { reason: String }, // Hardware resource errors GpuRequired { reason: String }, // GPU acceleration required } ``` ### Error Recovery The service implements sophisticated error recovery mechanisms: 1. **Automatic Retries**: Transient failures trigger automatic retries with exponential backoff 2. **Checkpoint Recovery**: Training resumes from last valid checkpoint on recoverable errors 3. **Resource Reallocation**: Automatic reallocation of resources on hardware failures 4. **Graceful Degradation**: CPU fallback when GPU resources unavailable 5. **Data Validation**: Comprehensive data validation with automatic cleaning ### Safety Controls Mathematical safety is enforced through multiple layers: ```rust pub struct GradientSafetyConfig { pub max_gradient_norm: f64, // Gradient clipping threshold pub nan_detection_enabled: bool, // NaN detection pub inf_detection_enabled: bool, // Infinity detection pub numerical_stability_threshold: f64, // Numerical stability threshold pub gradient_explosion_threshold: f64, // Gradient explosion detection } ``` ## Production Deployment ### Performance Requirements - **Training Latency**: < 10ms per forward pass for real-time training - **Memory Efficiency**: < 2GB memory usage per model during training - **GPU Utilization**: > 80% GPU utilization during active training - **Fault Tolerance**: Automatic recovery within 30 seconds of failures - **Scalability**: Support for 100+ concurrent training jobs ### Security & Compliance - **Authentication**: mTLS certificate-based authentication - **Authorization**: Role-based access control (RBAC) - **Audit Logging**: Comprehensive audit trail for all training activities - **Data Privacy**: PII anonymization and secure data handling - **Regulatory Compliance**: SOX, MiFID II compliance for financial models ### Infrastructure Requirements **Minimum Requirements:** - 2x NVIDIA A100 GPUs (40GB VRAM each) - 64 cores CPU (Intel Xeon or AMD EPYC) - 256GB RAM - 2TB NVMe SSD storage - 10GbE network connectivity **Recommended Production Setup:** - 8x NVIDIA H100 GPUs (80GB VRAM each) - 128 cores CPU - 1TB RAM - 10TB NVMe SSD storage - InfiniBand network (200Gb/s) - Redundant power supplies ### Monitoring & Alerting Production deployment includes comprehensive monitoring: ```yaml # Prometheus metrics configuration training_metrics: - name: ml_training_jobs_total help: Total number of training jobs labels: [model_type, status] - name: ml_training_duration_seconds help: Training duration histogram buckets: [60, 300, 900, 3600, 14400] - name: ml_gpu_utilization_percent help: GPU utilization percentage labels: [gpu_id, job_id] - name: ml_memory_usage_bytes help: Memory usage during training labels: [job_id, memory_type] ``` ### High Availability The service supports high availability deployment: - **Active-Passive Failover**: Automatic failover to standby instances - **Load Balancing**: Intelligent load balancing across GPU resources - **Data Replication**: Real-time data replication for disaster recovery - **Health Checks**: Comprehensive health monitoring with auto-restart - **Rolling Updates**: Zero-downtime updates and deployments --- ## Conclusion The MLTrainingService provides a production-ready, enterprise-grade machine learning training platform specifically designed for high-frequency trading applications. With comprehensive safety controls, real-time monitoring, and seamless integration with the Foxhunt ecosystem, it enables reliable and scalable ML model development and deployment. For additional information, see: - [TLI Operations Manual](./TLI_OPERATIONS_MANUAL.md) - [System Architecture](./SYSTEM_ARCHITECTURE.md) - [Performance Tuning](./PERFORMANCE_TUNING.md) - [Security Documentation](./SECURITY.md)