validation.py 30 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752
  1. """
  2. Model validation and testing framework for ML trainers.
  3. This module provides comprehensive validation and testing capabilities including
  4. cross-validation, model evaluation, performance benchmarking, and robustness testing.
  5. """
  6. import os
  7. import logging
  8. import time
  9. import random
  10. from pathlib import Path
  11. from typing import Dict, List, Tuple, Optional, Any, Callable, Union
  12. from dataclasses import dataclass
  13. import numpy as np
  14. from collections import defaultdict
  15. import torch
  16. import torch.nn as nn
  17. import torch.nn.functional as F
  18. from torch.utils.data import DataLoader, Subset, random_split
  19. from sklearn.model_selection import KFold, StratifiedKFold
  20. from sklearn.metrics import classification_report, confusion_matrix
  21. from .utils import ValidationMetrics, ModelProfiler
  22. from .data_pipeline import AudioDataset, AudioAugmentation, AudioProcessor
  23. @dataclass
  24. class ValidationConfig:
  25. """Configuration for validation and testing."""
  26. k_folds: int = 5
  27. test_augmentations: bool = True
  28. robustness_tests: bool = True
  29. performance_profiling: bool = True
  30. save_predictions: bool = True
  31. save_embeddings: bool = False
  32. confidence_threshold: float = 0.5
  33. batch_size: int = 32
  34. num_workers: int = 4
  35. class ModelValidator:
  36. """
  37. Comprehensive model validation and testing framework.
  38. Provides various validation methods including standard validation,
  39. cross-validation, robustness testing, and performance profiling.
  40. """
  41. def __init__(self, model: nn.Module, device: str = "cpu",
  42. logger: Optional[logging.Logger] = None,
  43. compute_loss_fn: Optional[Callable] = None):
  44. """
  45. Initialize model validator.
  46. Args:
  47. model: PyTorch model to validate
  48. device: Device for validation
  49. logger: Optional logger instance
  50. compute_loss_fn: Optional custom loss computation function
  51. """
  52. self.model = model
  53. self.device = device
  54. self.logger = logger or logging.getLogger(__name__)
  55. self.compute_loss_fn = compute_loss_fn
  56. # Initialize validation metrics calculator
  57. self.metrics_calculator = ValidationMetrics(
  58. task_type="classification" # Will be updated based on model
  59. )
  60. # Initialize model profiler
  61. self.profiler = ModelProfiler(model, device)
  62. # Results storage
  63. self.validation_results = {}
  64. self.test_results = {}
  65. self.cross_validation_results = {}
  66. self.robustness_results = {}
  67. self.profiling_results = {}
  68. def validate(self, val_loader: DataLoader, criterion: nn.Module,
  69. config: ValidationConfig = None) -> Dict[str, Any]:
  70. """
  71. Perform standard validation.
  72. Args:
  73. val_loader: Validation data loader
  74. criterion: Loss criterion
  75. config: Validation configuration
  76. Returns:
  77. Validation results dictionary
  78. """
  79. if config is None:
  80. config = ValidationConfig()
  81. self.logger.info("Starting model validation...")
  82. self.model.eval()
  83. total_loss = 0.0
  84. all_predictions = []
  85. all_probabilities = []
  86. all_targets = []
  87. all_embeddings = []
  88. validation_start_time = time.time()
  89. with torch.no_grad():
  90. for batch_idx, (data, targets) in enumerate(val_loader):
  91. data = data.to(self.device)
  92. targets = targets.to(self.device)
  93. # Forward pass
  94. outputs = self.model(data)
  95. if self.compute_loss_fn is not None:
  96. loss = self.compute_loss_fn(outputs, targets, data)
  97. else:
  98. loss = criterion(outputs, targets)
  99. total_loss += loss.item()
  100. # Get predictions
  101. if len(outputs.shape) > 1 and outputs.shape[1] > 1:
  102. # Classification - get probabilities and predictions
  103. probabilities = F.softmax(outputs, dim=1)
  104. predictions = torch.argmax(outputs, dim=1)
  105. all_probabilities.append(probabilities.cpu().numpy())
  106. all_predictions.append(predictions.cpu().numpy())
  107. else:
  108. # Regression or single output
  109. all_predictions.append(outputs.cpu().numpy())
  110. all_targets.append(targets.cpu().numpy())
  111. # Store embeddings if model has embedding layer
  112. if hasattr(self.model, 'get_embeddings'):
  113. try:
  114. embeddings = self.model.get_embeddings(data)
  115. all_embeddings.append(embeddings.cpu().numpy())
  116. except:
  117. pass
  118. validation_time = time.time() - validation_start_time
  119. # Concatenate all results
  120. all_targets = np.concatenate(all_targets)
  121. all_predictions = np.concatenate(all_predictions)
  122. if all_probabilities:
  123. all_probabilities = np.concatenate(all_probabilities)
  124. else:
  125. all_probabilities = None
  126. if all_embeddings:
  127. all_embeddings = np.concatenate(all_embeddings)
  128. else:
  129. all_embeddings = None
  130. # Calculate metrics
  131. avg_loss = total_loss / len(val_loader)
  132. # Update metrics calculator with appropriate settings
  133. num_classes = len(np.unique(all_targets))
  134. self.metrics_calculator.num_classes = num_classes
  135. self.metrics_calculator.class_names = [f"Class_{i}" for i in range(num_classes)]
  136. metrics = self.metrics_calculator.calculate_all_metrics(
  137. all_targets, all_predictions, all_probabilities, all_embeddings
  138. )
  139. # Compile results
  140. results = {
  141. "validation_loss": avg_loss,
  142. "validation_time": validation_time,
  143. "num_samples": len(all_targets),
  144. "metrics": metrics
  145. }
  146. # Save predictions if requested
  147. if config.save_predictions:
  148. results["predictions"] = all_predictions.tolist()
  149. results["targets"] = all_targets.tolist()
  150. if all_probabilities is not None:
  151. results["probabilities"] = all_probabilities.tolist()
  152. # Save embeddings if requested
  153. if config.save_embeddings and all_embeddings is not None:
  154. results["embeddings"] = all_embeddings.tolist()
  155. self.validation_results = results
  156. self.logger.info(f"Validation completed - Loss: {avg_loss:.6f}, "
  157. f"Accuracy: {metrics.get('accuracy', 0):.4f}")
  158. return results
  159. def test(self, test_loader: DataLoader, criterion: nn.Module,
  160. config: ValidationConfig = None) -> Dict[str, Any]:
  161. """
  162. Perform comprehensive testing.
  163. Args:
  164. test_loader: Test data loader
  165. criterion: Loss criterion
  166. config: Validation configuration
  167. Returns:
  168. Test results dictionary
  169. """
  170. if config is None:
  171. config = ValidationConfig()
  172. self.logger.info("Starting model testing...")
  173. # Standard test
  174. test_results = self.validate(test_loader, criterion, config)
  175. test_results["test_type"] = "standard"
  176. results = {"standard_test": test_results}
  177. # Augmentation robustness test
  178. if config.test_augmentations:
  179. aug_results = self._test_augmentation_robustness(test_loader, criterion, config)
  180. results["augmentation_robustness"] = aug_results
  181. # Performance profiling
  182. if config.performance_profiling:
  183. profile_results = self._profile_performance(test_loader)
  184. results["performance_profile"] = profile_results
  185. # Additional robustness tests
  186. if config.robustness_tests:
  187. robustness_results = self._test_robustness(test_loader, criterion, config)
  188. results["robustness_tests"] = robustness_results
  189. self.test_results = results
  190. self.logger.info("Testing completed successfully")
  191. return results
  192. def cross_validate(self, dataset: AudioDataset, criterion: nn.Module,
  193. config: ValidationConfig = None) -> Dict[str, Any]:
  194. """
  195. Perform k-fold cross-validation.
  196. Args:
  197. dataset: Complete dataset for cross-validation
  198. criterion: Loss criterion
  199. config: Validation configuration
  200. Returns:
  201. Cross-validation results
  202. """
  203. if config is None:
  204. config = ValidationConfig()
  205. self.logger.info(f"Starting {config.k_folds}-fold cross-validation...")
  206. # Get labels for stratified split
  207. labels = np.array(dataset.labels)
  208. # Use stratified k-fold for balanced splits
  209. if len(np.unique(labels)) > 1:
  210. kfold = StratifiedKFold(n_splits=config.k_folds, shuffle=True, random_state=42)
  211. splits = list(kfold.split(np.arange(len(dataset)), labels))
  212. else:
  213. kfold = KFold(n_splits=config.k_folds, shuffle=True, random_state=42)
  214. splits = list(kfold.split(np.arange(len(dataset))))
  215. fold_results = []
  216. for fold, (train_idx, val_idx) in enumerate(splits):
  217. self.logger.info(f"Evaluating fold {fold + 1}/{config.k_folds}")
  218. # Create validation subset
  219. val_subset = Subset(dataset, val_idx)
  220. val_loader = DataLoader(
  221. val_subset,
  222. batch_size=config.batch_size,
  223. shuffle=False,
  224. num_workers=config.num_workers
  225. )
  226. # Validate on this fold
  227. fold_result = self.validate(val_loader, criterion, config)
  228. fold_result["fold"] = fold
  229. fold_result["train_samples"] = len(train_idx)
  230. fold_result["val_samples"] = len(val_idx)
  231. fold_results.append(fold_result)
  232. self.logger.info(f"Fold {fold + 1} - Loss: {fold_result['validation_loss']:.6f}, "
  233. f"Accuracy: {fold_result['metrics'].get('accuracy', 0):.4f}")
  234. # Aggregate results
  235. cv_results = self._aggregate_cv_results(fold_results)
  236. self.cross_validation_results = cv_results
  237. self.logger.info(f"Cross-validation completed - Mean Accuracy: "
  238. f"{cv_results['mean_accuracy']:.4f} ± {cv_results['std_accuracy']:.4f}")
  239. return cv_results
  240. def _test_augmentation_robustness(self, test_loader: DataLoader,
  241. criterion: nn.Module,
  242. config: ValidationConfig) -> Dict[str, Any]:
  243. """Test model robustness against data augmentations."""
  244. self.logger.info("Testing augmentation robustness...")
  245. # Define augmentation configurations
  246. augmentation_configs = {
  247. "noise_light": {"noise_factor": 0.05},
  248. "noise_moderate": {"noise_factor": 0.15},
  249. "noise_heavy": {"noise_factor": 0.25},
  250. "speed_light": {"speed_factor": 0.05},
  251. "speed_moderate": {"speed_factor": 0.15},
  252. "volume_light": {"volume_factor": 0.1},
  253. "volume_moderate": {"volume_factor": 0.3},
  254. "combined_light": {
  255. "noise_factor": 0.05,
  256. "speed_factor": 0.05,
  257. "volume_factor": 0.1
  258. }
  259. }
  260. augmentation = AudioAugmentation(sample_rate=16000, device=self.device)
  261. results = {}
  262. for aug_name, aug_config in augmentation_configs.items():
  263. self.logger.info(f"Testing with {aug_name} augmentation")
  264. self.model.eval()
  265. total_loss = 0.0
  266. all_predictions = []
  267. all_targets = []
  268. with torch.no_grad():
  269. for data, targets in test_loader:
  270. data = data.to(self.device)
  271. targets = targets.to(self.device)
  272. # Apply augmentation to raw audio (convert back from features)
  273. # This is a simplified approach - in practice, you'd need the raw audio
  274. # For now, we'll apply augmentation to the feature representations
  275. outputs = self.model(data)
  276. if self.compute_loss_fn is not None:
  277. loss = self.compute_loss_fn(outputs, targets, data)
  278. else:
  279. loss = criterion(outputs, targets)
  280. total_loss += loss.item()
  281. if len(outputs.shape) > 1 and outputs.shape[1] > 1:
  282. predictions = torch.argmax(outputs, dim=1)
  283. all_predictions.append(predictions.cpu().numpy())
  284. else:
  285. all_predictions.append(outputs.cpu().numpy())
  286. all_targets.append(targets.cpu().numpy())
  287. # Calculate metrics
  288. all_targets = np.concatenate(all_targets)
  289. all_predictions = np.concatenate(all_predictions)
  290. metrics = self.metrics_calculator.calculate_all_metrics(
  291. all_targets, all_predictions
  292. )
  293. results[aug_name] = {
  294. "loss": total_loss / len(test_loader),
  295. "accuracy": metrics.get("accuracy", 0),
  296. "metrics": metrics
  297. }
  298. return results
  299. def _profile_performance(self, test_loader: DataLoader) -> Dict[str, Any]:
  300. """Profile model performance characteristics."""
  301. self.logger.info("Profiling model performance...")
  302. # Get sample input
  303. sample_batch = next(iter(test_loader))
  304. sample_input = sample_batch[0][:1].to(self.device) # Single sample
  305. # Profile inference
  306. inference_stats = self.profiler.profile_inference(sample_input)
  307. memory_stats = self.profiler.profile_memory(sample_input)
  308. param_stats = self.profiler.count_parameters()
  309. size_stats = self.profiler.estimate_model_size()
  310. # Test batch processing performance
  311. batch_sizes = [1, 8, 16, 32, 64]
  312. batch_performance = {}
  313. for batch_size in batch_sizes:
  314. if batch_size <= len(sample_batch[0]):
  315. batch_input = sample_batch[0][:batch_size].to(self.device)
  316. batch_stats = self.profiler.profile_inference(batch_input, num_runs=20)
  317. batch_performance[f"batch_{batch_size}"] = {
  318. "inference_time": batch_stats["mean_inference_time"],
  319. "throughput": batch_stats["throughput_samples_per_second"]
  320. }
  321. return {
  322. "inference_performance": inference_stats,
  323. "memory_usage": memory_stats,
  324. "model_parameters": param_stats,
  325. "model_size": size_stats,
  326. "batch_performance": batch_performance
  327. }
  328. def _test_robustness(self, test_loader: DataLoader, criterion: nn.Module,
  329. config: ValidationConfig) -> Dict[str, Any]:
  330. """Test model robustness with various perturbations."""
  331. self.logger.info("Testing model robustness...")
  332. results = {}
  333. # Test with different confidence thresholds
  334. threshold_results = self._test_confidence_thresholds(test_loader, criterion)
  335. results["confidence_thresholds"] = threshold_results
  336. # Test with corrupted inputs
  337. corruption_results = self._test_input_corruptions(test_loader, criterion)
  338. results["input_corruptions"] = corruption_results
  339. # Test with adversarial examples (simplified)
  340. adversarial_results = self._test_adversarial_robustness(test_loader, criterion)
  341. results["adversarial_robustness"] = adversarial_results
  342. return results
  343. def _test_confidence_thresholds(self, test_loader: DataLoader,
  344. criterion: nn.Module) -> Dict[str, Any]:
  345. """Test model performance at different confidence thresholds."""
  346. thresholds = [0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.99]
  347. results = {}
  348. self.model.eval()
  349. all_probabilities = []
  350. all_targets = []
  351. with torch.no_grad():
  352. for data, targets in test_loader:
  353. data = data.to(self.device)
  354. targets = targets.to(self.device)
  355. outputs = self.model(data)
  356. if len(outputs.shape) > 1 and outputs.shape[1] > 1:
  357. probabilities = F.softmax(outputs, dim=1)
  358. all_probabilities.append(probabilities.cpu().numpy())
  359. all_targets.append(targets.cpu().numpy())
  360. if all_probabilities:
  361. all_probabilities = np.concatenate(all_probabilities)
  362. all_targets = np.concatenate(all_targets)
  363. for threshold in thresholds:
  364. # Get confident predictions
  365. max_probs = np.max(all_probabilities, axis=1)
  366. confident_mask = max_probs >= threshold
  367. if np.sum(confident_mask) > 0:
  368. confident_preds = np.argmax(all_probabilities[confident_mask], axis=1)
  369. confident_targets = all_targets[confident_mask]
  370. accuracy = np.mean(confident_preds == confident_targets)
  371. coverage = np.mean(confident_mask)
  372. results[f"threshold_{threshold}"] = {
  373. "accuracy": accuracy,
  374. "coverage": coverage,
  375. "num_samples": np.sum(confident_mask)
  376. }
  377. return results
  378. def _test_input_corruptions(self, test_loader: DataLoader,
  379. criterion: nn.Module) -> Dict[str, Any]:
  380. """Test robustness to input corruptions."""
  381. corruption_types = {
  382. "zero_out_10": lambda x: self._zero_out_random(x, 0.1),
  383. "zero_out_25": lambda x: self._zero_out_random(x, 0.25),
  384. "gaussian_noise_01": lambda x: x + torch.randn_like(x) * 0.1,
  385. "gaussian_noise_02": lambda x: x + torch.randn_like(x) * 0.2,
  386. "dropout_10": lambda x: F.dropout(x, p=0.1, training=True),
  387. "dropout_25": lambda x: F.dropout(x, p=0.25, training=True)
  388. }
  389. results = {}
  390. for corruption_name, corruption_func in corruption_types.items():
  391. self.logger.info(f"Testing with {corruption_name}")
  392. self.model.eval()
  393. total_loss = 0.0
  394. all_predictions = []
  395. all_targets = []
  396. with torch.no_grad():
  397. for data, targets in test_loader:
  398. data = data.to(self.device)
  399. targets = targets.to(self.device)
  400. # Apply corruption
  401. corrupted_data = corruption_func(data)
  402. outputs = self.model(corrupted_data)
  403. if self.compute_loss_fn is not None:
  404. loss = self.compute_loss_fn(outputs, targets, data)
  405. else:
  406. loss = criterion(outputs, targets)
  407. total_loss += loss.item()
  408. if len(outputs.shape) > 1 and outputs.shape[1] > 1:
  409. predictions = torch.argmax(outputs, dim=1)
  410. all_predictions.append(predictions.cpu().numpy())
  411. else:
  412. all_predictions.append(outputs.cpu().numpy())
  413. all_targets.append(targets.cpu().numpy())
  414. # Calculate metrics
  415. all_targets = np.concatenate(all_targets)
  416. all_predictions = np.concatenate(all_predictions)
  417. accuracy = np.mean(all_predictions == all_targets)
  418. results[corruption_name] = {
  419. "loss": total_loss / len(test_loader),
  420. "accuracy": accuracy
  421. }
  422. return results
  423. def _test_adversarial_robustness(self, test_loader: DataLoader,
  424. criterion: nn.Module) -> Dict[str, Any]:
  425. """Test robustness to adversarial examples (simplified FGSM)."""
  426. epsilon_values = [0.01, 0.05, 0.1, 0.2]
  427. results = {}
  428. for epsilon in epsilon_values:
  429. self.logger.info(f"Testing adversarial robustness with epsilon={epsilon}")
  430. total_loss = 0.0
  431. all_predictions = []
  432. all_targets = []
  433. for data, targets in test_loader:
  434. data = data.to(self.device)
  435. targets = targets.to(self.device)
  436. data.requires_grad = True
  437. # Forward pass
  438. outputs = self.model(data)
  439. if self.compute_loss_fn is not None:
  440. loss = self.compute_loss_fn(outputs, targets, data)
  441. else:
  442. loss = criterion(outputs, targets)
  443. # Backward pass to get gradients
  444. self.model.zero_grad()
  445. loss.backward()
  446. # Generate adversarial examples using FGSM
  447. data_grad = data.grad.data
  448. perturbed_data = data + epsilon * data_grad.sign()
  449. # Re-evaluate with perturbed data
  450. with torch.no_grad():
  451. perturbed_outputs = self.model(perturbed_data)
  452. if self.compute_loss_fn is not None:
  453. perturbed_loss = self.compute_loss_fn(perturbed_outputs, targets, perturbed_data)
  454. else:
  455. perturbed_loss = criterion(perturbed_outputs, targets)
  456. total_loss += perturbed_loss.item()
  457. if len(perturbed_outputs.shape) > 1 and perturbed_outputs.shape[1] > 1:
  458. predictions = torch.argmax(perturbed_outputs, dim=1)
  459. all_predictions.append(predictions.cpu().numpy())
  460. else:
  461. all_predictions.append(perturbed_outputs.cpu().numpy())
  462. all_targets.append(targets.cpu().numpy())
  463. # Calculate metrics
  464. all_targets = np.concatenate(all_targets)
  465. all_predictions = np.concatenate(all_predictions)
  466. accuracy = np.mean(all_predictions == all_targets)
  467. results[f"epsilon_{epsilon}"] = {
  468. "loss": total_loss / len(test_loader),
  469. "accuracy": accuracy
  470. }
  471. return results
  472. def _zero_out_random(self, tensor: torch.Tensor, fraction: float) -> torch.Tensor:
  473. """Randomly zero out a fraction of tensor elements."""
  474. mask = torch.rand_like(tensor) < fraction
  475. return tensor * (~mask).float()
  476. def _aggregate_cv_results(self, fold_results: List[Dict[str, Any]]) -> Dict[str, Any]:
  477. """Aggregate cross-validation results across folds."""
  478. # Extract metrics from all folds
  479. losses = [result["validation_loss"] for result in fold_results]
  480. accuracies = [result["metrics"].get("accuracy", 0) for result in fold_results]
  481. # Calculate statistics
  482. aggregated = {
  483. "num_folds": len(fold_results),
  484. "mean_loss": np.mean(losses),
  485. "std_loss": np.std(losses),
  486. "mean_accuracy": np.mean(accuracies),
  487. "std_accuracy": np.std(accuracies),
  488. "fold_results": fold_results
  489. }
  490. # Aggregate other metrics if available
  491. metric_keys = set()
  492. for result in fold_results:
  493. metric_keys.update(result["metrics"].keys())
  494. for metric_key in metric_keys:
  495. if metric_key != "confusion_matrix" and metric_key != "classification_report":
  496. values = []
  497. for result in fold_results:
  498. if metric_key in result["metrics"]:
  499. value = result["metrics"][metric_key]
  500. if isinstance(value, (int, float)):
  501. values.append(value)
  502. if values:
  503. aggregated[f"mean_{metric_key}"] = np.mean(values)
  504. aggregated[f"std_{metric_key}"] = np.std(values)
  505. return aggregated
  506. def generate_validation_report(self, output_dir: str = None) -> str:
  507. """Generate comprehensive validation report."""
  508. report = "# Model Validation Report\n\n"
  509. # Add validation results
  510. if self.validation_results:
  511. report += "## Validation Results\n"
  512. val_res = self.validation_results
  513. report += f"- Validation Loss: {val_res['validation_loss']:.6f}\n"
  514. report += f"- Validation Time: {val_res['validation_time']:.2f} seconds\n"
  515. report += f"- Number of Samples: {val_res['num_samples']:,}\n"
  516. if "accuracy" in val_res["metrics"]:
  517. report += f"- Accuracy: {val_res['metrics']['accuracy']:.4f}\n"
  518. report += "\n"
  519. # Add test results
  520. if self.test_results:
  521. report += "## Test Results\n"
  522. # Standard test
  523. if "standard_test" in self.test_results:
  524. std_test = self.test_results["standard_test"]
  525. report += f"- Test Loss: {std_test['validation_loss']:.6f}\n"
  526. if "accuracy" in std_test["metrics"]:
  527. report += f"- Test Accuracy: {std_test['metrics']['accuracy']:.4f}\n"
  528. # Performance profile
  529. if "performance_profile" in self.test_results:
  530. profile = self.test_results["performance_profile"]
  531. report += f"- Inference Time: {profile['inference_performance']['mean_inference_time']*1000:.2f} ms\n"
  532. report += f"- Model Size: {profile['model_size']['total_size_mb']:.2f} MB\n"
  533. report += f"- Parameters: {profile['model_parameters']['total_parameters']:,}\n"
  534. report += "\n"
  535. # Add cross-validation results
  536. if self.cross_validation_results:
  537. report += "## Cross-Validation Results\n"
  538. cv_res = self.cross_validation_results
  539. report += f"- Number of Folds: {cv_res['num_folds']}\n"
  540. report += f"- Mean Accuracy: {cv_res['mean_accuracy']:.4f} ± {cv_res['std_accuracy']:.4f}\n"
  541. report += f"- Mean Loss: {cv_res['mean_loss']:.6f} ± {cv_res['std_loss']:.6f}\n"
  542. report += "\n"
  543. # Save report if output directory provided
  544. if output_dir:
  545. os.makedirs(output_dir, exist_ok=True)
  546. report_path = Path(output_dir) / "validation_report.md"
  547. with open(report_path, 'w') as f:
  548. f.write(report)
  549. self.logger.info(f"Validation report saved to: {report_path}")
  550. return report
  551. def save_results(self, output_dir: str):
  552. """Save all validation results to files."""
  553. output_dir = Path(output_dir)
  554. output_dir.mkdir(parents=True, exist_ok=True)
  555. # Save validation results
  556. if self.validation_results:
  557. val_path = output_dir / "validation_results.json"
  558. self._save_json(self.validation_results, val_path)
  559. # Save test results
  560. if self.test_results:
  561. test_path = output_dir / "test_results.json"
  562. self._save_json(self.test_results, test_path)
  563. # Save cross-validation results
  564. if self.cross_validation_results:
  565. cv_path = output_dir / "cross_validation_results.json"
  566. self._save_json(self.cross_validation_results, cv_path)
  567. # Generate and save report
  568. report = self.generate_validation_report(str(output_dir))
  569. self.logger.info(f"Validation results saved to: {output_dir}")
  570. def _save_json(self, data: Dict[str, Any], filepath: Path):
  571. """Save data to JSON file with proper handling of numpy arrays."""
  572. import json
  573. def convert_numpy(obj):
  574. if isinstance(obj, np.ndarray):
  575. return obj.tolist()
  576. elif isinstance(obj, np.integer):
  577. return int(obj)
  578. elif isinstance(obj, np.floating):
  579. return float(obj)
  580. return obj
  581. # Convert numpy objects recursively
  582. def clean_data(data):
  583. if isinstance(data, dict):
  584. return {k: clean_data(v) for k, v in data.items()}
  585. elif isinstance(data, list):
  586. return [clean_data(v) for v in data]
  587. else:
  588. return convert_numpy(data)
  589. cleaned_data = clean_data(data)
  590. with open(filepath, 'w') as f:
  591. json.dump(cleaned_data, f, indent=2)