Implement comprehensive Runpod deployment with S3 volume mount architecture for FP32 ML model training on Tesla V100 GPUs. ## Infrastructure Components ### Deployment Scripts (scripts/) - runpod_deploy.sh: Master deployment orchestrator (8-step workflow) - runpod_upload.sh: S3 upload for binaries and test data - upload_env_to_runpod.sh: Secure .env credentials upload - runpod_deploy_test.sh: Prerequisites validation ### Docker Configuration - Dockerfile.runpod: Multi-stage CUDA 12.1 runtime (~2GB, no binaries) - entrypoint.sh: Volume verification and training execution - Architecture: Volume mount (NO S3 downloads in pods) ### S3 Configuration - Bucket: se3zdnb5o4 (Iceland region: eur-is-1) - Endpoint: https://s3api-eur-is-1.runpod.io - Structure: binaries/, test_data/, models/, .env ### OpenTofu Infrastructure (terraform/runpod/) - main.tf: Pod and volume resources - variables.tf: Configuration variables - outputs.tf: Pod connection info - Security: NO credentials in state (uses volume .env) ## Deployment Assets Uploaded ### Training Binaries (77MB) - train_tft_parquet (23M) - TFT-225 features - train_mamba2_parquet (22M) - MAMBA-2 state space - train_dqn (22M) - Deep Q-Network - train_ppo (13M) - Proximal Policy Optimization ### Test Data (13.8 MB) - 9 Parquet files: ES.FUT, NQ.FUT, 6E.FUT, ZN.FUT (180-day datasets) ### Credentials - .env file (1.5 KB, private access, chmod 600) ## Documentation ### Deployment Guides - RUNPOD_DEPLOYMENT_READY_SUMMARY.md: Complete deployment status - RUNPOD_VOLUME_DEPLOYMENT_GUIDE.md: Step-by-step guide (42KB) - RUNPOD_DEPLOYMENT_QUICK_START.md: Quick reference - RUNPOD_UPLOAD_GUIDE.md: S3 upload instructions - RUNPOD_VOLUME_CONFIGURATION_COMPLETE.md: S3 setup report - RUNPOD_S3_PARQUET_UPLOAD_REPORT.md: Data upload verification ### Architecture Documentation - RUNPOD_VOLUME_MOUNT_ARCHITECTURE.md: Volume mount design - RUNPOD_S3_ARCHITECTURE_DIAGRAM.txt: S3 API vs filesystem access - DOCKERFILE_RUNPOD_FINAL_SUMMARY.md: Docker image specification ### Decision Documentation - RUNPOD_DEPLOYMENT_CHECKLIST.md: Go/no-go decision matrix (27KB) - RUNPOD_DEPLOYMENT_DECISION_TREE.md: Decision workflow - FP32_RUNPOD_DEPLOYMENT_READY.md: FP32 deployment readiness ## QAT Enhancements ### Core QAT Infrastructure - ml/src/memory_optimization/qat.rs: Enhanced QAT observer (+226 lines) - ml/src/memory_optimization/auto_batch_size.rs: OOM recovery (+84 lines) - ml/src/tft/qat_tft.rs: QAT TFT wrapper (+154 lines) - ml/src/trainers/tft.rs: QAT training integration (+433 lines) - ml/src/qat_metrics_exporter.rs: NEW - QAT metrics export ### QAT Testing - ml/tests/qat_integration_tests.rs: NEW - Integration test suite - ml/tests/qat_gradient_clipping_test.rs: NEW - Gradient clipping tests - ml/tests/qat_device_consistency_test.rs: Device mismatch tests (+205 lines) - ml/tests/qat_accuracy_validation_test.rs: Accuracy validation - ml/tests/qat_tft_integration_test.rs: TFT QAT integration ### QAT Documentation - ml/docs/QAT_GUIDE.md: Comprehensive QAT guide (+616 lines) - ml/docs/QAT_GRADIENT_CHECKPOINTING_WORKAROUND.md: NEW - Workaround guide - QAT_BLOCKERS_ROOT_CAUSE_ANALYSIS.md: P0 blocker analysis (44KB) - QAT_ACCURACY_VALIDATION_REPORT.md: Accuracy comparison - QAT_GRADIENT_CLIPPING_VALIDATION_REPORT.md: Clipping validation ### QAT Monitoring - config/grafana/dashboards/qat-training-metrics.json: NEW - Grafana dashboard ## AWS CLI Configuration ### Credentials Setup - ~/.aws/credentials: Runpod profile configured - Access Key: user_2xxA3XcIFj16yfL3aBon9niiSpr - Secret Key: (from RUNPOD_S3_SECRET) - ~/.aws/config: Iceland region (eur-is-1) ## Production Readiness ### FP32 Models: ✅ READY FOR DEPLOYMENT - DQN: 15-20s training, ~6MB GPU memory - PPO: 7-10s training, ~145MB GPU memory - MAMBA-2: 2-3 min training, ~164MB GPU memory - TFT-225: 3-5 min training, ~500MB GPU memory - Total GPU Budget: 815MB (fits on 4GB+ Tesla V100) ### QAT Models: 🔴 BLOCKED - 24 tests implemented but DO NOT COMPILE (11 errors) - 3 P0 blockers: device mismatch, gradient checkpointing, OOM recovery - Timeline: 1-2 weeks to fix (13h P0 fixes + validation) ### Wave D Features: ✅ OPERATIONAL - 225 features fully integrated - Feature extraction: 5.10μs/bar (196x faster than target) - Wave D backtest: Sharpe 2.00, Win Rate 60%, Drawdown 15% - Database migration 045: Applied cleanly, zero conflicts ## Cost Analysis ### One-Time Setup - Network Volume: $4/month (50GB SSD) - Upload costs: FREE (S3 API included) ### Per Training Run (TFT-225) - GPU: Tesla V100-PCIE-16GB @ $0.29/hr - Training Time: ~4 hours - Cost per run: $1.16 ### Monthly (20 Training Runs) - Storage: $4.00/month - Training: $23.20/month (20 runs × $1.16) - Total: $27.20/month ## Security ### Credentials Management - ✅ NO credentials in Docker image - ✅ NO credentials in Terraform state - ✅ .env gitignored and not committed - ✅ .env file private on S3 (HTTP 401 on public access) - ✅ Docker Hub repository PRIVATE (jgrusewski/foxhunt) ### Access Control - S3 API: Local client uploads only - Volume mount: Pod filesystem access only - Authentication: AWS CLI with Runpod profile required ## Next Steps 1. ✅ COMPLETE: Build Docker image 2. ⏳ PENDING: Push to Docker Hub 3. ⏳ PENDING: Deploy pod via Runpod console 4. ⏳ PENDING: Validate training on Tesla V100 ## Performance Targets - Build time: 5-10 min - Upload time: ~20 sec (90MB total) - Pod startup: ~30 sec - Training time: 3-5 min (TFT-225) - Total deployment: ~40 min from start to first training run ## Test Status - FP32 tests: 597/608 passing (98.2%) - QAT tests: 0/24 passing (compilation errors) - Overall: 2,062/2,086 passing (98.8% excluding QAT) 🤖 Generated with Claude Code (https://claude.com/claude-code) Co-Authored-By: Claude <noreply@anthropic.com>
547 lines
19 KiB
Rust
547 lines
19 KiB
Rust
//! Model Registry Checkpoint Integration Tests
|
|
//!
|
|
//! TDD tests for checkpoint versioning, metadata tracking, and production model registration.
|
|
//! Wave 10 Agent 10.8 - Training → Paper Trading Integration
|
|
|
|
use chrono::Utc;
|
|
use ml::model_registry::{ModelRegistry, ModelVersionMetadata};
|
|
use ml::{MLResult, ModelType};
|
|
use std::path::PathBuf;
|
|
|
|
// Test database URL
|
|
const TEST_DB_URL: &str = "postgresql://foxhunt:foxhunt_dev_password@localhost:5432/foxhunt";
|
|
const TEST_S3_PATH: &str = "s3://foxhunt-ml-models-test/";
|
|
|
|
/// Test 1: Register trained DQN model with checkpoint path
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_register_dqn_checkpoint() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
// Find latest DQN checkpoint
|
|
let checkpoint_path = PathBuf::from(
|
|
"/home/jgrusewski/Work/foxhunt/ml/trained_models/production/dqn/dqn_epoch_30.safetensors",
|
|
);
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"dqn-production-v1.0.0".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://foxhunt-ml-models/dqn/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add checkpoint metadata
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(30));
|
|
metadata.add_hyperparameter("batch_size", serde_json::json!(128));
|
|
metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001));
|
|
|
|
metadata.add_metric("final_loss", serde_json::json!(0.0342));
|
|
metadata.add_metric("validation_loss", serde_json::json!(0.0356));
|
|
|
|
metadata.add_metadata(
|
|
"checkpoint_path",
|
|
checkpoint_path.to_string_lossy().to_string(),
|
|
);
|
|
metadata.add_metadata("training_duration_hours", "2.5".to_string());
|
|
|
|
metadata.set_checksum("sha256:dqn_epoch_30_checksum".to_string());
|
|
|
|
// Register
|
|
registry.register_version(&metadata).await?;
|
|
|
|
// Verify retrieval
|
|
let retrieved = registry
|
|
.get_model_by_version("dqn-production-v1.0.0")
|
|
.await?;
|
|
assert_eq!(retrieved.model_id, "dqn-production-v1.0.0");
|
|
assert_eq!(retrieved.model_type, ModelType::DQN);
|
|
assert_eq!(retrieved.version, "1.0.0");
|
|
assert!(retrieved.metadata.contains_key("checkpoint_path"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 2: Register trained PPO model with actor-critic checkpoints
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_register_ppo_checkpoint() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let actor_checkpoint = PathBuf::from("/home/jgrusewski/Work/foxhunt/ml/trained_models/production/ppo/ppo_actor_epoch_420.safetensors");
|
|
let critic_checkpoint = PathBuf::from("/home/jgrusewski/Work/foxhunt/ml/trained_models/production/ppo/ppo_critic_epoch_420.safetensors");
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"ppo-production-v1.0.0".to_string(),
|
|
ModelType::PPO,
|
|
"1.0.0".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://foxhunt-ml-models/ppo/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add PPO-specific hyperparameters
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(420));
|
|
metadata.add_hyperparameter("batch_size", serde_json::json!(64));
|
|
metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0003));
|
|
metadata.add_hyperparameter("gamma", serde_json::json!(0.99));
|
|
metadata.add_hyperparameter("gae_lambda", serde_json::json!(0.95));
|
|
|
|
metadata.add_metric("final_actor_loss", serde_json::json!(0.0152));
|
|
metadata.add_metric("final_critic_loss", serde_json::json!(0.0089));
|
|
metadata.add_metric("avg_reward", serde_json::json!(45.3));
|
|
|
|
metadata.add_metadata(
|
|
"actor_checkpoint_path",
|
|
actor_checkpoint.to_string_lossy().to_string(),
|
|
);
|
|
metadata.add_metadata(
|
|
"critic_checkpoint_path",
|
|
critic_checkpoint.to_string_lossy().to_string(),
|
|
);
|
|
|
|
metadata.set_checksum("sha256:ppo_epoch_420_checksum".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("ppo-production-v1.0.0")
|
|
.await?;
|
|
assert_eq!(retrieved.model_type, ModelType::PPO);
|
|
assert!(retrieved.metadata.contains_key("actor_checkpoint_path"));
|
|
assert!(retrieved.metadata.contains_key("critic_checkpoint_path"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 3: Register trained MAMBA-2 model with training metrics
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_register_mamba2_checkpoint() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"mamba2-production-v1.0.0".to_string(),
|
|
ModelType::MAMBA,
|
|
"1.0.0".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://foxhunt-ml-models/mamba2/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add MAMBA-2 hyperparameters
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(24));
|
|
metadata.add_hyperparameter("batch_size", serde_json::json!(32));
|
|
metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001));
|
|
metadata.add_hyperparameter("d_model", serde_json::json!(256));
|
|
metadata.add_hyperparameter("n_layers", serde_json::json!(6));
|
|
metadata.add_hyperparameter("state_size", serde_json::json!(16));
|
|
|
|
metadata.add_metric("best_val_loss", serde_json::json!(1.4318895660848898));
|
|
metadata.add_metric("best_epoch", serde_json::json!(3));
|
|
metadata.add_metric("final_perplexity", serde_json::json!(4.1866025848353));
|
|
|
|
metadata.add_metadata(
|
|
"checkpoint_path",
|
|
"/home/jgrusewski/Work/foxhunt/ml/checkpoints/mamba2_dbn/".to_string(),
|
|
);
|
|
metadata.add_metadata("training_duration_hours", "0.031".to_string());
|
|
|
|
metadata.set_checksum("sha256:mamba2_epoch_24_checksum".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("mamba2-production-v1.0.0")
|
|
.await?;
|
|
assert_eq!(retrieved.model_type, ModelType::MAMBA);
|
|
|
|
// Verify metrics
|
|
let metrics = retrieved.metrics.as_object().unwrap();
|
|
assert!(metrics.contains_key("best_val_loss"));
|
|
assert!(metrics.contains_key("final_perplexity"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 4: Register trained TFT model with multiple checkpoints
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_register_tft_checkpoint() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let checkpoint_path = PathBuf::from(
|
|
"/home/jgrusewski/Work/foxhunt/ml/trained_models/production/tft/tft_epoch_100.safetensors",
|
|
);
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"tft-production-v1.0.0".to_string(),
|
|
ModelType::TFT,
|
|
"1.0.0".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://foxhunt-ml-models/tft/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add TFT hyperparameters
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(100));
|
|
metadata.add_hyperparameter("batch_size", serde_json::json!(256));
|
|
metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001));
|
|
metadata.add_hyperparameter("hidden_size", serde_json::json!(256));
|
|
metadata.add_hyperparameter("num_attention_heads", serde_json::json!(8));
|
|
|
|
metadata.add_metric("final_loss", serde_json::json!(0.0198));
|
|
metadata.add_metric("validation_loss", serde_json::json!(0.0213));
|
|
metadata.add_metric("sharpe_ratio", serde_json::json!(2.4));
|
|
|
|
metadata.add_metadata(
|
|
"checkpoint_path",
|
|
checkpoint_path.to_string_lossy().to_string(),
|
|
);
|
|
|
|
metadata.set_checksum("sha256:tft_epoch_100_checksum".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("tft-production-v1.0.0")
|
|
.await?;
|
|
assert_eq!(retrieved.model_type, ModelType::TFT);
|
|
|
|
// Verify hyperparameters
|
|
let hyperparams = retrieved.hyperparameters.as_object().unwrap();
|
|
assert_eq!(hyperparams.get("epochs").unwrap(), &serde_json::json!(100));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 5: Register TFT-INT8 quantized model
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_register_tft_int8_checkpoint() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"tft-int8-production-v1.0.0".to_string(),
|
|
ModelType::TFT,
|
|
"1.0.0-int8".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://foxhunt-ml-models/tft-int8/1.0.0/".to_string(),
|
|
);
|
|
|
|
metadata.add_hyperparameter("quantization", serde_json::json!("int8"));
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(100));
|
|
|
|
metadata.add_metric("inference_latency_ms", serde_json::json!(3.2));
|
|
metadata.add_metric("model_size_mb", serde_json::json!(128));
|
|
|
|
metadata.add_metadata("quantization_method", "static_int8".to_string());
|
|
metadata.add_metadata("optimization_level", "production".to_string());
|
|
|
|
metadata.set_checksum("sha256:tft_int8_checksum".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("tft-int8-production-v1.0.0")
|
|
.await?;
|
|
assert_eq!(retrieved.version, "1.0.0-int8");
|
|
assert!(retrieved.metadata.contains_key("quantization_method"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 6: Version increment handling
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_version_increment() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
// Register v1.0.0
|
|
let mut metadata_v1 = ModelVersionMetadata::new(
|
|
"dqn-version-test-v1.0.0".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/1.0.0/".to_string(),
|
|
);
|
|
metadata_v1.add_metric("loss", serde_json::json!(0.05));
|
|
registry.register_version(&metadata_v1).await?;
|
|
|
|
// Register v1.1.0 (improvement)
|
|
let mut metadata_v1_1 = ModelVersionMetadata::new(
|
|
"dqn-version-test-v1.1.0".to_string(),
|
|
ModelType::DQN,
|
|
"1.1.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/1.1.0/".to_string(),
|
|
);
|
|
metadata_v1_1.add_metric("loss", serde_json::json!(0.03));
|
|
registry.register_version(&metadata_v1_1).await?;
|
|
|
|
// Register v2.0.0 (major update)
|
|
let mut metadata_v2 = ModelVersionMetadata::new(
|
|
"dqn-version-test-v2.0.0".to_string(),
|
|
ModelType::DQN,
|
|
"2.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/2.0.0/".to_string(),
|
|
);
|
|
metadata_v2.add_metric("loss", serde_json::json!(0.01));
|
|
registry.register_version(&metadata_v2).await?;
|
|
|
|
// Verify all versions exist
|
|
let v1 = registry
|
|
.get_model_by_version("dqn-version-test-v1.0.0")
|
|
.await?;
|
|
assert_eq!(v1.version, "1.0.0");
|
|
|
|
let v1_1 = registry
|
|
.get_model_by_version("dqn-version-test-v1.1.0")
|
|
.await?;
|
|
assert_eq!(v1_1.version, "1.1.0");
|
|
|
|
let v2 = registry
|
|
.get_model_by_version("dqn-version-test-v2.0.0")
|
|
.await?;
|
|
assert_eq!(v2.version, "2.0.0");
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 7: Checkpoint path validation
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_checkpoint_path_metadata() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let checkpoint_path = PathBuf::from(
|
|
"/home/jgrusewski/Work/foxhunt/ml/trained_models/production/dqn/dqn_epoch_30.safetensors",
|
|
);
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"dqn-checkpoint-path-test".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/1.0.0/".to_string(),
|
|
);
|
|
|
|
metadata.add_metadata(
|
|
"checkpoint_path",
|
|
checkpoint_path.to_string_lossy().to_string(),
|
|
);
|
|
metadata.add_metadata("checkpoint_format", "safetensors".to_string());
|
|
metadata.add_metadata("checkpoint_size_mb", "256".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("dqn-checkpoint-path-test")
|
|
.await?;
|
|
|
|
assert!(retrieved.metadata.contains_key("checkpoint_path"));
|
|
assert_eq!(
|
|
retrieved.metadata.get("checkpoint_format").unwrap(),
|
|
"safetensors"
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 8: Multi-model registry query
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_multi_model_registry_query() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
// Register multiple models
|
|
let model_types = vec![
|
|
(ModelType::DQN, "dqn-multi-test"),
|
|
(ModelType::PPO, "ppo-multi-test"),
|
|
(ModelType::MAMBA, "mamba-multi-test"),
|
|
(ModelType::TFT, "tft-multi-test"),
|
|
];
|
|
|
|
for (model_type, model_id) in model_types {
|
|
let metadata = ModelVersionMetadata::new(
|
|
model_id.to_string(),
|
|
model_type,
|
|
"1.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
format!("s3://test/{}/1.0.0/", model_id),
|
|
);
|
|
registry.register_version(&metadata).await?;
|
|
}
|
|
|
|
// Query by type
|
|
let dqn_models = registry.get_models_by_type(ModelType::DQN).await?;
|
|
assert!(dqn_models.iter().any(|m| m.model_id == "dqn-multi-test"));
|
|
|
|
let ppo_models = registry.get_models_by_type(ModelType::PPO).await?;
|
|
assert!(ppo_models.iter().any(|m| m.model_id == "ppo-multi-test"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 9: Production model promotion workflow
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_production_promotion_workflow() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let metadata = ModelVersionMetadata::new(
|
|
"dqn-promotion-test".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Start as experimental
|
|
assert!(metadata.is_experimental);
|
|
assert!(!metadata.is_production);
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
// Promote to production
|
|
registry.mark_production("dqn-promotion-test").await?;
|
|
|
|
// Verify production status
|
|
let retrieved = registry.get_model_by_version("dqn-promotion-test").await?;
|
|
assert!(retrieved.is_production);
|
|
assert!(!retrieved.is_experimental);
|
|
|
|
// Verify in production query
|
|
let production_models = registry.get_production_models().await?;
|
|
assert!(production_models
|
|
.iter()
|
|
.any(|m| m.model_id == "dqn-promotion-test"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 10: Training metrics metadata
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_training_metrics_metadata() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"dqn-metrics-test".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"test_data".to_string(),
|
|
"s3://test/dqn/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add comprehensive metrics
|
|
metadata.add_metric("final_loss", serde_json::json!(0.0342));
|
|
metadata.add_metric("validation_loss", serde_json::json!(0.0356));
|
|
metadata.add_metric("best_epoch", serde_json::json!(28));
|
|
metadata.add_metric("total_epochs", serde_json::json!(30));
|
|
metadata.add_metric("training_duration_hours", serde_json::json!(2.5));
|
|
metadata.add_metric("gpu_memory_used_gb", serde_json::json!(3.2));
|
|
metadata.add_metric("avg_epoch_time_seconds", serde_json::json!(300));
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry.get_model_by_version("dqn-metrics-test").await?;
|
|
|
|
let metrics = retrieved.metrics.as_object().unwrap();
|
|
assert_eq!(
|
|
metrics.get("final_loss").unwrap(),
|
|
&serde_json::json!(0.0342)
|
|
);
|
|
assert_eq!(metrics.get("best_epoch").unwrap(), &serde_json::json!(28));
|
|
assert!(metrics.contains_key("gpu_memory_used_gb"));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 11: List all checkpoints for a model type
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_list_checkpoints_by_type() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
// Register multiple DQN checkpoints
|
|
for epoch in [10, 20, 30] {
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
format!("dqn-checkpoint-epoch-{}", epoch),
|
|
ModelType::DQN,
|
|
format!("1.0.{}", epoch),
|
|
"test_data".to_string(),
|
|
format!("s3://test/dqn/1.0.{}/", epoch),
|
|
);
|
|
metadata.add_metadata("epoch", epoch.to_string());
|
|
registry.register_version(&metadata).await?;
|
|
}
|
|
|
|
let dqn_models = registry.get_models_by_type(ModelType::DQN).await?;
|
|
let checkpoint_models: Vec<_> = dqn_models
|
|
.iter()
|
|
.filter(|m| m.model_id.starts_with("dqn-checkpoint-epoch-"))
|
|
.collect();
|
|
|
|
assert!(checkpoint_models.len() >= 3);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Test 12: Checkpoint metadata completeness
|
|
#[tokio::test]
|
|
#[ignore = "Requires PostgreSQL"]
|
|
async fn test_checkpoint_metadata_completeness() -> MLResult<()> {
|
|
let registry = ModelRegistry::new(TEST_DB_URL, TEST_S3_PATH).await?;
|
|
|
|
let mut metadata = ModelVersionMetadata::new(
|
|
"complete-metadata-test".to_string(),
|
|
ModelType::DQN,
|
|
"1.0.0".to_string(),
|
|
"ES.FUT_2024_Q4".to_string(),
|
|
"s3://test/dqn/1.0.0/".to_string(),
|
|
);
|
|
|
|
// Add comprehensive metadata
|
|
metadata.add_hyperparameter("epochs", serde_json::json!(30));
|
|
metadata.add_hyperparameter("batch_size", serde_json::json!(128));
|
|
metadata.add_hyperparameter("learning_rate", serde_json::json!(0.0001));
|
|
metadata.add_hyperparameter("gamma", serde_json::json!(0.99));
|
|
metadata.add_hyperparameter("epsilon_start", serde_json::json!(1.0));
|
|
metadata.add_hyperparameter("epsilon_end", serde_json::json!(0.01));
|
|
|
|
metadata.add_metric("final_loss", serde_json::json!(0.0342));
|
|
metadata.add_metric("validation_loss", serde_json::json!(0.0356));
|
|
metadata.add_metric("sharpe_ratio", serde_json::json!(2.1));
|
|
metadata.add_metric("max_drawdown", serde_json::json!(0.12));
|
|
|
|
metadata.add_metadata(
|
|
"checkpoint_path",
|
|
"/path/to/checkpoint.safetensors".to_string(),
|
|
);
|
|
metadata.add_metadata("training_date", Utc::now().to_rfc3339());
|
|
metadata.add_metadata("cuda_version", "12.1".to_string());
|
|
metadata.add_metadata("pytorch_version", "2.0.0".to_string());
|
|
|
|
metadata.set_checksum("sha256:complete_metadata_checksum".to_string());
|
|
|
|
registry.register_version(&metadata).await?;
|
|
|
|
let retrieved = registry
|
|
.get_model_by_version("complete-metadata-test")
|
|
.await?;
|
|
|
|
// Verify hyperparameters
|
|
let hyperparams = retrieved.hyperparameters.as_object().unwrap();
|
|
assert_eq!(hyperparams.len(), 6);
|
|
|
|
// Verify metrics
|
|
let metrics = retrieved.metrics.as_object().unwrap();
|
|
assert_eq!(metrics.len(), 4);
|
|
|
|
// Verify metadata
|
|
assert_eq!(retrieved.metadata.len(), 4);
|
|
assert!(retrieved.metadata.contains_key("checkpoint_path"));
|
|
assert!(retrieved.metadata.contains_key("cuda_version"));
|
|
|
|
Ok(())
|
|
}
|