| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752 |
- """
- Model validation and testing framework for ML trainers.
- This module provides comprehensive validation and testing capabilities including
- cross-validation, model evaluation, performance benchmarking, and robustness testing.
- """
- import os
- import logging
- import time
- import random
- from pathlib import Path
- from typing import Dict, List, Tuple, Optional, Any, Callable, Union
- from dataclasses import dataclass
- import numpy as np
- from collections import defaultdict
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- from torch.utils.data import DataLoader, Subset, random_split
- from sklearn.model_selection import KFold, StratifiedKFold
- from sklearn.metrics import classification_report, confusion_matrix
- from .utils import ValidationMetrics, ModelProfiler
- from .data_pipeline import AudioDataset, AudioAugmentation, AudioProcessor
- @dataclass
- class ValidationConfig:
- """Configuration for validation and testing."""
- k_folds: int = 5
- test_augmentations: bool = True
- robustness_tests: bool = True
- performance_profiling: bool = True
- save_predictions: bool = True
- save_embeddings: bool = False
- confidence_threshold: float = 0.5
- batch_size: int = 32
- num_workers: int = 4
- class ModelValidator:
- """
- Comprehensive model validation and testing framework.
-
- Provides various validation methods including standard validation,
- cross-validation, robustness testing, and performance profiling.
- """
-
- def __init__(self, model: nn.Module, device: str = "cpu",
- logger: Optional[logging.Logger] = None,
- compute_loss_fn: Optional[Callable] = None):
- """
- Initialize model validator.
-
- Args:
- model: PyTorch model to validate
- device: Device for validation
- logger: Optional logger instance
- compute_loss_fn: Optional custom loss computation function
- """
- self.model = model
- self.device = device
- self.logger = logger or logging.getLogger(__name__)
- self.compute_loss_fn = compute_loss_fn
-
- # Initialize validation metrics calculator
- self.metrics_calculator = ValidationMetrics(
- task_type="classification" # Will be updated based on model
- )
-
- # Initialize model profiler
- self.profiler = ModelProfiler(model, device)
-
- # Results storage
- self.validation_results = {}
- self.test_results = {}
- self.cross_validation_results = {}
- self.robustness_results = {}
- self.profiling_results = {}
-
- def validate(self, val_loader: DataLoader, criterion: nn.Module,
- config: ValidationConfig = None) -> Dict[str, Any]:
- """
- Perform standard validation.
-
- Args:
- val_loader: Validation data loader
- criterion: Loss criterion
- config: Validation configuration
-
- Returns:
- Validation results dictionary
- """
- if config is None:
- config = ValidationConfig()
-
- self.logger.info("Starting model validation...")
-
- self.model.eval()
- total_loss = 0.0
- all_predictions = []
- all_probabilities = []
- all_targets = []
- all_embeddings = []
-
- validation_start_time = time.time()
-
- with torch.no_grad():
- for batch_idx, (data, targets) in enumerate(val_loader):
- data = data.to(self.device)
- targets = targets.to(self.device)
-
- # Forward pass
- outputs = self.model(data)
- if self.compute_loss_fn is not None:
- loss = self.compute_loss_fn(outputs, targets, data)
- else:
- loss = criterion(outputs, targets)
-
- total_loss += loss.item()
-
- # Get predictions
- if len(outputs.shape) > 1 and outputs.shape[1] > 1:
- # Classification - get probabilities and predictions
- probabilities = F.softmax(outputs, dim=1)
- predictions = torch.argmax(outputs, dim=1)
-
- all_probabilities.append(probabilities.cpu().numpy())
- all_predictions.append(predictions.cpu().numpy())
- else:
- # Regression or single output
- all_predictions.append(outputs.cpu().numpy())
-
- all_targets.append(targets.cpu().numpy())
-
- # Store embeddings if model has embedding layer
- if hasattr(self.model, 'get_embeddings'):
- try:
- embeddings = self.model.get_embeddings(data)
- all_embeddings.append(embeddings.cpu().numpy())
- except:
- pass
-
- validation_time = time.time() - validation_start_time
-
- # Concatenate all results
- all_targets = np.concatenate(all_targets)
- all_predictions = np.concatenate(all_predictions)
-
- if all_probabilities:
- all_probabilities = np.concatenate(all_probabilities)
- else:
- all_probabilities = None
-
- if all_embeddings:
- all_embeddings = np.concatenate(all_embeddings)
- else:
- all_embeddings = None
-
- # Calculate metrics
- avg_loss = total_loss / len(val_loader)
-
- # Update metrics calculator with appropriate settings
- num_classes = len(np.unique(all_targets))
- self.metrics_calculator.num_classes = num_classes
- self.metrics_calculator.class_names = [f"Class_{i}" for i in range(num_classes)]
-
- metrics = self.metrics_calculator.calculate_all_metrics(
- all_targets, all_predictions, all_probabilities, all_embeddings
- )
-
- # Compile results
- results = {
- "validation_loss": avg_loss,
- "validation_time": validation_time,
- "num_samples": len(all_targets),
- "metrics": metrics
- }
-
- # Save predictions if requested
- if config.save_predictions:
- results["predictions"] = all_predictions.tolist()
- results["targets"] = all_targets.tolist()
- if all_probabilities is not None:
- results["probabilities"] = all_probabilities.tolist()
-
- # Save embeddings if requested
- if config.save_embeddings and all_embeddings is not None:
- results["embeddings"] = all_embeddings.tolist()
-
- self.validation_results = results
- self.logger.info(f"Validation completed - Loss: {avg_loss:.6f}, "
- f"Accuracy: {metrics.get('accuracy', 0):.4f}")
-
- return results
-
- def test(self, test_loader: DataLoader, criterion: nn.Module,
- config: ValidationConfig = None) -> Dict[str, Any]:
- """
- Perform comprehensive testing.
-
- Args:
- test_loader: Test data loader
- criterion: Loss criterion
- config: Validation configuration
-
- Returns:
- Test results dictionary
- """
- if config is None:
- config = ValidationConfig()
-
- self.logger.info("Starting model testing...")
-
- # Standard test
- test_results = self.validate(test_loader, criterion, config)
- test_results["test_type"] = "standard"
-
- results = {"standard_test": test_results}
-
- # Augmentation robustness test
- if config.test_augmentations:
- aug_results = self._test_augmentation_robustness(test_loader, criterion, config)
- results["augmentation_robustness"] = aug_results
-
- # Performance profiling
- if config.performance_profiling:
- profile_results = self._profile_performance(test_loader)
- results["performance_profile"] = profile_results
-
- # Additional robustness tests
- if config.robustness_tests:
- robustness_results = self._test_robustness(test_loader, criterion, config)
- results["robustness_tests"] = robustness_results
-
- self.test_results = results
- self.logger.info("Testing completed successfully")
-
- return results
-
- def cross_validate(self, dataset: AudioDataset, criterion: nn.Module,
- config: ValidationConfig = None) -> Dict[str, Any]:
- """
- Perform k-fold cross-validation.
-
- Args:
- dataset: Complete dataset for cross-validation
- criterion: Loss criterion
- config: Validation configuration
-
- Returns:
- Cross-validation results
- """
- if config is None:
- config = ValidationConfig()
-
- self.logger.info(f"Starting {config.k_folds}-fold cross-validation...")
-
- # Get labels for stratified split
- labels = np.array(dataset.labels)
-
- # Use stratified k-fold for balanced splits
- if len(np.unique(labels)) > 1:
- kfold = StratifiedKFold(n_splits=config.k_folds, shuffle=True, random_state=42)
- splits = list(kfold.split(np.arange(len(dataset)), labels))
- else:
- kfold = KFold(n_splits=config.k_folds, shuffle=True, random_state=42)
- splits = list(kfold.split(np.arange(len(dataset))))
-
- fold_results = []
-
- for fold, (train_idx, val_idx) in enumerate(splits):
- self.logger.info(f"Evaluating fold {fold + 1}/{config.k_folds}")
-
- # Create validation subset
- val_subset = Subset(dataset, val_idx)
- val_loader = DataLoader(
- val_subset,
- batch_size=config.batch_size,
- shuffle=False,
- num_workers=config.num_workers
- )
-
- # Validate on this fold
- fold_result = self.validate(val_loader, criterion, config)
- fold_result["fold"] = fold
- fold_result["train_samples"] = len(train_idx)
- fold_result["val_samples"] = len(val_idx)
-
- fold_results.append(fold_result)
-
- self.logger.info(f"Fold {fold + 1} - Loss: {fold_result['validation_loss']:.6f}, "
- f"Accuracy: {fold_result['metrics'].get('accuracy', 0):.4f}")
-
- # Aggregate results
- cv_results = self._aggregate_cv_results(fold_results)
-
- self.cross_validation_results = cv_results
- self.logger.info(f"Cross-validation completed - Mean Accuracy: "
- f"{cv_results['mean_accuracy']:.4f} ± {cv_results['std_accuracy']:.4f}")
-
- return cv_results
-
- def _test_augmentation_robustness(self, test_loader: DataLoader,
- criterion: nn.Module,
- config: ValidationConfig) -> Dict[str, Any]:
- """Test model robustness against data augmentations."""
- self.logger.info("Testing augmentation robustness...")
-
- # Define augmentation configurations
- augmentation_configs = {
- "noise_light": {"noise_factor": 0.05},
- "noise_moderate": {"noise_factor": 0.15},
- "noise_heavy": {"noise_factor": 0.25},
- "speed_light": {"speed_factor": 0.05},
- "speed_moderate": {"speed_factor": 0.15},
- "volume_light": {"volume_factor": 0.1},
- "volume_moderate": {"volume_factor": 0.3},
- "combined_light": {
- "noise_factor": 0.05,
- "speed_factor": 0.05,
- "volume_factor": 0.1
- }
- }
-
- augmentation = AudioAugmentation(sample_rate=16000, device=self.device)
- results = {}
-
- for aug_name, aug_config in augmentation_configs.items():
- self.logger.info(f"Testing with {aug_name} augmentation")
-
- self.model.eval()
- total_loss = 0.0
- all_predictions = []
- all_targets = []
-
- with torch.no_grad():
- for data, targets in test_loader:
- data = data.to(self.device)
- targets = targets.to(self.device)
-
- # Apply augmentation to raw audio (convert back from features)
- # This is a simplified approach - in practice, you'd need the raw audio
- # For now, we'll apply augmentation to the feature representations
-
- outputs = self.model(data)
- if self.compute_loss_fn is not None:
- loss = self.compute_loss_fn(outputs, targets, data)
- else:
- loss = criterion(outputs, targets)
-
- total_loss += loss.item()
-
- if len(outputs.shape) > 1 and outputs.shape[1] > 1:
- predictions = torch.argmax(outputs, dim=1)
- all_predictions.append(predictions.cpu().numpy())
- else:
- all_predictions.append(outputs.cpu().numpy())
-
- all_targets.append(targets.cpu().numpy())
-
- # Calculate metrics
- all_targets = np.concatenate(all_targets)
- all_predictions = np.concatenate(all_predictions)
-
- metrics = self.metrics_calculator.calculate_all_metrics(
- all_targets, all_predictions
- )
-
- results[aug_name] = {
- "loss": total_loss / len(test_loader),
- "accuracy": metrics.get("accuracy", 0),
- "metrics": metrics
- }
-
- return results
-
- def _profile_performance(self, test_loader: DataLoader) -> Dict[str, Any]:
- """Profile model performance characteristics."""
- self.logger.info("Profiling model performance...")
-
- # Get sample input
- sample_batch = next(iter(test_loader))
- sample_input = sample_batch[0][:1].to(self.device) # Single sample
-
- # Profile inference
- inference_stats = self.profiler.profile_inference(sample_input)
- memory_stats = self.profiler.profile_memory(sample_input)
- param_stats = self.profiler.count_parameters()
- size_stats = self.profiler.estimate_model_size()
-
- # Test batch processing performance
- batch_sizes = [1, 8, 16, 32, 64]
- batch_performance = {}
-
- for batch_size in batch_sizes:
- if batch_size <= len(sample_batch[0]):
- batch_input = sample_batch[0][:batch_size].to(self.device)
- batch_stats = self.profiler.profile_inference(batch_input, num_runs=20)
- batch_performance[f"batch_{batch_size}"] = {
- "inference_time": batch_stats["mean_inference_time"],
- "throughput": batch_stats["throughput_samples_per_second"]
- }
-
- return {
- "inference_performance": inference_stats,
- "memory_usage": memory_stats,
- "model_parameters": param_stats,
- "model_size": size_stats,
- "batch_performance": batch_performance
- }
-
- def _test_robustness(self, test_loader: DataLoader, criterion: nn.Module,
- config: ValidationConfig) -> Dict[str, Any]:
- """Test model robustness with various perturbations."""
- self.logger.info("Testing model robustness...")
-
- results = {}
-
- # Test with different confidence thresholds
- threshold_results = self._test_confidence_thresholds(test_loader, criterion)
- results["confidence_thresholds"] = threshold_results
-
- # Test with corrupted inputs
- corruption_results = self._test_input_corruptions(test_loader, criterion)
- results["input_corruptions"] = corruption_results
-
- # Test with adversarial examples (simplified)
- adversarial_results = self._test_adversarial_robustness(test_loader, criterion)
- results["adversarial_robustness"] = adversarial_results
-
- return results
-
- def _test_confidence_thresholds(self, test_loader: DataLoader,
- criterion: nn.Module) -> Dict[str, Any]:
- """Test model performance at different confidence thresholds."""
- thresholds = [0.5, 0.6, 0.7, 0.8, 0.9, 0.95, 0.99]
- results = {}
-
- self.model.eval()
- all_probabilities = []
- all_targets = []
-
- with torch.no_grad():
- for data, targets in test_loader:
- data = data.to(self.device)
- targets = targets.to(self.device)
-
- outputs = self.model(data)
-
- if len(outputs.shape) > 1 and outputs.shape[1] > 1:
- probabilities = F.softmax(outputs, dim=1)
- all_probabilities.append(probabilities.cpu().numpy())
- all_targets.append(targets.cpu().numpy())
-
- if all_probabilities:
- all_probabilities = np.concatenate(all_probabilities)
- all_targets = np.concatenate(all_targets)
-
- for threshold in thresholds:
- # Get confident predictions
- max_probs = np.max(all_probabilities, axis=1)
- confident_mask = max_probs >= threshold
-
- if np.sum(confident_mask) > 0:
- confident_preds = np.argmax(all_probabilities[confident_mask], axis=1)
- confident_targets = all_targets[confident_mask]
-
- accuracy = np.mean(confident_preds == confident_targets)
- coverage = np.mean(confident_mask)
-
- results[f"threshold_{threshold}"] = {
- "accuracy": accuracy,
- "coverage": coverage,
- "num_samples": np.sum(confident_mask)
- }
-
- return results
-
- def _test_input_corruptions(self, test_loader: DataLoader,
- criterion: nn.Module) -> Dict[str, Any]:
- """Test robustness to input corruptions."""
- corruption_types = {
- "zero_out_10": lambda x: self._zero_out_random(x, 0.1),
- "zero_out_25": lambda x: self._zero_out_random(x, 0.25),
- "gaussian_noise_01": lambda x: x + torch.randn_like(x) * 0.1,
- "gaussian_noise_02": lambda x: x + torch.randn_like(x) * 0.2,
- "dropout_10": lambda x: F.dropout(x, p=0.1, training=True),
- "dropout_25": lambda x: F.dropout(x, p=0.25, training=True)
- }
-
- results = {}
-
- for corruption_name, corruption_func in corruption_types.items():
- self.logger.info(f"Testing with {corruption_name}")
-
- self.model.eval()
- total_loss = 0.0
- all_predictions = []
- all_targets = []
-
- with torch.no_grad():
- for data, targets in test_loader:
- data = data.to(self.device)
- targets = targets.to(self.device)
-
- # Apply corruption
- corrupted_data = corruption_func(data)
-
- outputs = self.model(corrupted_data)
- if self.compute_loss_fn is not None:
- loss = self.compute_loss_fn(outputs, targets, data)
- else:
- loss = criterion(outputs, targets)
-
- total_loss += loss.item()
-
- if len(outputs.shape) > 1 and outputs.shape[1] > 1:
- predictions = torch.argmax(outputs, dim=1)
- all_predictions.append(predictions.cpu().numpy())
- else:
- all_predictions.append(outputs.cpu().numpy())
-
- all_targets.append(targets.cpu().numpy())
-
- # Calculate metrics
- all_targets = np.concatenate(all_targets)
- all_predictions = np.concatenate(all_predictions)
-
- accuracy = np.mean(all_predictions == all_targets)
-
- results[corruption_name] = {
- "loss": total_loss / len(test_loader),
- "accuracy": accuracy
- }
-
- return results
-
- def _test_adversarial_robustness(self, test_loader: DataLoader,
- criterion: nn.Module) -> Dict[str, Any]:
- """Test robustness to adversarial examples (simplified FGSM)."""
- epsilon_values = [0.01, 0.05, 0.1, 0.2]
- results = {}
-
- for epsilon in epsilon_values:
- self.logger.info(f"Testing adversarial robustness with epsilon={epsilon}")
-
- total_loss = 0.0
- all_predictions = []
- all_targets = []
-
- for data, targets in test_loader:
- data = data.to(self.device)
- targets = targets.to(self.device)
- data.requires_grad = True
-
- # Forward pass
- outputs = self.model(data)
- if self.compute_loss_fn is not None:
- loss = self.compute_loss_fn(outputs, targets, data)
- else:
- loss = criterion(outputs, targets)
-
- # Backward pass to get gradients
- self.model.zero_grad()
- loss.backward()
-
- # Generate adversarial examples using FGSM
- data_grad = data.grad.data
- perturbed_data = data + epsilon * data_grad.sign()
-
- # Re-evaluate with perturbed data
- with torch.no_grad():
- perturbed_outputs = self.model(perturbed_data)
- if self.compute_loss_fn is not None:
- perturbed_loss = self.compute_loss_fn(perturbed_outputs, targets, perturbed_data)
- else:
- perturbed_loss = criterion(perturbed_outputs, targets)
-
- total_loss += perturbed_loss.item()
-
- if len(perturbed_outputs.shape) > 1 and perturbed_outputs.shape[1] > 1:
- predictions = torch.argmax(perturbed_outputs, dim=1)
- all_predictions.append(predictions.cpu().numpy())
- else:
- all_predictions.append(perturbed_outputs.cpu().numpy())
-
- all_targets.append(targets.cpu().numpy())
-
- # Calculate metrics
- all_targets = np.concatenate(all_targets)
- all_predictions = np.concatenate(all_predictions)
-
- accuracy = np.mean(all_predictions == all_targets)
-
- results[f"epsilon_{epsilon}"] = {
- "loss": total_loss / len(test_loader),
- "accuracy": accuracy
- }
-
- return results
-
- def _zero_out_random(self, tensor: torch.Tensor, fraction: float) -> torch.Tensor:
- """Randomly zero out a fraction of tensor elements."""
- mask = torch.rand_like(tensor) < fraction
- return tensor * (~mask).float()
-
- def _aggregate_cv_results(self, fold_results: List[Dict[str, Any]]) -> Dict[str, Any]:
- """Aggregate cross-validation results across folds."""
- # Extract metrics from all folds
- losses = [result["validation_loss"] for result in fold_results]
- accuracies = [result["metrics"].get("accuracy", 0) for result in fold_results]
-
- # Calculate statistics
- aggregated = {
- "num_folds": len(fold_results),
- "mean_loss": np.mean(losses),
- "std_loss": np.std(losses),
- "mean_accuracy": np.mean(accuracies),
- "std_accuracy": np.std(accuracies),
- "fold_results": fold_results
- }
-
- # Aggregate other metrics if available
- metric_keys = set()
- for result in fold_results:
- metric_keys.update(result["metrics"].keys())
-
- for metric_key in metric_keys:
- if metric_key != "confusion_matrix" and metric_key != "classification_report":
- values = []
- for result in fold_results:
- if metric_key in result["metrics"]:
- value = result["metrics"][metric_key]
- if isinstance(value, (int, float)):
- values.append(value)
-
- if values:
- aggregated[f"mean_{metric_key}"] = np.mean(values)
- aggregated[f"std_{metric_key}"] = np.std(values)
-
- return aggregated
-
- def generate_validation_report(self, output_dir: str = None) -> str:
- """Generate comprehensive validation report."""
- report = "# Model Validation Report\n\n"
-
- # Add validation results
- if self.validation_results:
- report += "## Validation Results\n"
- val_res = self.validation_results
- report += f"- Validation Loss: {val_res['validation_loss']:.6f}\n"
- report += f"- Validation Time: {val_res['validation_time']:.2f} seconds\n"
- report += f"- Number of Samples: {val_res['num_samples']:,}\n"
-
- if "accuracy" in val_res["metrics"]:
- report += f"- Accuracy: {val_res['metrics']['accuracy']:.4f}\n"
-
- report += "\n"
-
- # Add test results
- if self.test_results:
- report += "## Test Results\n"
-
- # Standard test
- if "standard_test" in self.test_results:
- std_test = self.test_results["standard_test"]
- report += f"- Test Loss: {std_test['validation_loss']:.6f}\n"
- if "accuracy" in std_test["metrics"]:
- report += f"- Test Accuracy: {std_test['metrics']['accuracy']:.4f}\n"
-
- # Performance profile
- if "performance_profile" in self.test_results:
- profile = self.test_results["performance_profile"]
- report += f"- Inference Time: {profile['inference_performance']['mean_inference_time']*1000:.2f} ms\n"
- report += f"- Model Size: {profile['model_size']['total_size_mb']:.2f} MB\n"
- report += f"- Parameters: {profile['model_parameters']['total_parameters']:,}\n"
-
- report += "\n"
-
- # Add cross-validation results
- if self.cross_validation_results:
- report += "## Cross-Validation Results\n"
- cv_res = self.cross_validation_results
- report += f"- Number of Folds: {cv_res['num_folds']}\n"
- report += f"- Mean Accuracy: {cv_res['mean_accuracy']:.4f} ± {cv_res['std_accuracy']:.4f}\n"
- report += f"- Mean Loss: {cv_res['mean_loss']:.6f} ± {cv_res['std_loss']:.6f}\n"
- report += "\n"
-
- # Save report if output directory provided
- if output_dir:
- os.makedirs(output_dir, exist_ok=True)
- report_path = Path(output_dir) / "validation_report.md"
- with open(report_path, 'w') as f:
- f.write(report)
- self.logger.info(f"Validation report saved to: {report_path}")
-
- return report
-
- def save_results(self, output_dir: str):
- """Save all validation results to files."""
- output_dir = Path(output_dir)
- output_dir.mkdir(parents=True, exist_ok=True)
-
- # Save validation results
- if self.validation_results:
- val_path = output_dir / "validation_results.json"
- self._save_json(self.validation_results, val_path)
-
- # Save test results
- if self.test_results:
- test_path = output_dir / "test_results.json"
- self._save_json(self.test_results, test_path)
-
- # Save cross-validation results
- if self.cross_validation_results:
- cv_path = output_dir / "cross_validation_results.json"
- self._save_json(self.cross_validation_results, cv_path)
-
- # Generate and save report
- report = self.generate_validation_report(str(output_dir))
-
- self.logger.info(f"Validation results saved to: {output_dir}")
-
- def _save_json(self, data: Dict[str, Any], filepath: Path):
- """Save data to JSON file with proper handling of numpy arrays."""
- import json
-
- def convert_numpy(obj):
- if isinstance(obj, np.ndarray):
- return obj.tolist()
- elif isinstance(obj, np.integer):
- return int(obj)
- elif isinstance(obj, np.floating):
- return float(obj)
- return obj
-
- # Convert numpy objects recursively
- def clean_data(data):
- if isinstance(data, dict):
- return {k: clean_data(v) for k, v in data.items()}
- elif isinstance(data, list):
- return [clean_data(v) for v in data]
- else:
- return convert_numpy(data)
-
- cleaned_data = clean_data(data)
-
- with open(filepath, 'w') as f:
- json.dump(cleaned_data, f, indent=2)
|