ASesYusuf1/SESA_Audio_Separation
14
1# coding: utf-82__author__ = 'PyTorch Optimization Benchmark Tool'3 4import argparse5import time6import torch7import numpy as np8from utils import get_model_from_config9from pytorch_backend import (10 PyTorchBackend, 11 PyTorchOptimizer, 12 benchmark_pytorch_optimizations,13 get_model_info14)15import sys16 17 18def load_checkpoint(checkpoint_path: str, model, device: str):19 """Load model from checkpoint."""20 print(f"Loading checkpoint from: {checkpoint_path}")21 22 checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)23 24 # Handle different checkpoint formats25 if isinstance(checkpoint, dict):26 if 'state_dict' in checkpoint:27 state_dict = checkpoint['state_dict']28 elif 'model' in checkpoint:29 state_dict = checkpoint['model']30 elif 'state' in checkpoint:31 state_dict = checkpoint['state']32 else:33 state_dict = checkpoint34 else:35 state_dict = checkpoint36 37 model.load_state_dict(state_dict, strict=False)38 model = model.eval().to(device)39 40 print("✓ Checkpoint loaded successfully")41 return model42 43 44def benchmark_optimization_modes(args):45 """46 Benchmark different PyTorch optimization modes.47 """48 parser = argparse.ArgumentParser(description="Benchmark PyTorch Optimization Modes")49 parser.add_argument("--model_type", type=str, required=True, help="Model type")50 parser.add_argument("--config_path", type=str, required=True, help="Config path")51 parser.add_argument("--start_check_point", type=str, required=True, help="Checkpoint path (.ckpt)")52 parser.add_argument("--device", type=str, default='cuda:0', help="Device")53 parser.add_argument("--num_iterations", type=int, default=100, help="Number of benchmark iterations")54 parser.add_argument("--warmup_iterations", type=int, default=10, help="Number of warmup iterations")55 parser.add_argument("--chunk_size", type=int, default=None, help="Override chunk size (optional)")56 parser.add_argument("--batch_size", type=int, default=1, help="Batch size")57 58 if args is None:59 args = parser.parse_args()60 else:61 args = parser.parse_args(args)62 63 # Check device64 if args.device.startswith('cuda') and not torch.cuda.is_available():65 print("❌ CUDA is not available!")66 return67 68 print("="*60)69 print("PyTorch Optimization Benchmark Tool")70 print("="*60)71 print(f"Model Type: {args.model_type}")72 print(f"Checkpoint: {args.start_check_point}")73 print(f"Device: {args.device}")74 print(f"Iterations: {args.num_iterations}")75 print("="*60)76 77 # Load model78 print("\n📦 Loading model...")79 model, config = get_model_from_config(args.model_type, args.config_path)80 model = load_checkpoint(args.start_check_point, model, args.device)81 82 # Get model info83 model_info = get_model_info(model)84 print(f"\n📊 Model Information:")85 print(f" Total Parameters: {model_info['total_parameters']:,}")86 print(f" Trainable Parameters: {model_info['trainable_parameters']:,}")87 print(f" Model Size: {model_info['model_size_mb']:.2f} MB")88 print(f" Device: {model_info['device']}")89 print(f" Dtype: {model_info['dtype']}")90 91 # Get chunk size92 if args.chunk_size:93 chunk_size = args.chunk_size94 else:95 chunk_size = config.audio.chunk_size96 97 num_channels = 298 input_shape = (args.batch_size, num_channels, chunk_size)99 100 print(f"\n📊 Test Configuration:")101 print(f" Batch Size: {args.batch_size}")102 print(f" Channels: {num_channels}")103 print(f" Chunk Size: {chunk_size}")104 print(f" Input Shape: {input_shape}")105 106 # Benchmark different optimization modes107 print("\n" + "="*60)108 print("Benchmarking Optimization Modes")109 print("="*60)110 111 results = benchmark_pytorch_optimizations(112 model=model,113 input_shape=input_shape,114 device=args.device,115 num_iterations=args.num_iterations,116 warmup_iterations=args.warmup_iterations117 )118 119 # Display results120 print("\n" + "="*60)121 print("📈 Benchmark Results")122 print("="*60)123 124 baseline = None125 for mode, time_ms in results.items():126 if time_ms is not None:127 if baseline is None:128 baseline = time_ms129 speedup = baseline / time_ms if time_ms > 0 else 0130 improvement = ((baseline - time_ms) / baseline) * 100 if baseline > 0 else 0131 132 print(f"\n{mode.upper()}:")133 print(f" Average Time: {time_ms:.2f} ms")134 print(f" Speedup: {speedup:.2f}x")135 print(f" Improvement: {improvement:.1f}%")136 137 print("\n" + "="*60)138 139 # Recommendations140 print("\n💡 Recommendations:")141 142 if results.get('compile') and results['compile'] < results['default']:143 print(" ✓ Use 'compile' mode for best performance (PyTorch 2.0+)")144 elif results.get('channels_last') and results['channels_last'] < results['default']:145 print(" ✓ Use 'channels_last' mode for better performance")146 else:147 print(" ✓ Default mode is optimal for your configuration")148 149 if args.device.startswith('cuda'):150 print(" ✓ Enable TF32 for Ampere GPUs (RTX 30xx+)")151 print(" ✓ Enable cuDNN benchmark for consistent input sizes")152 153 print("\n✅ Benchmark completed!")154 155 156def test_optimization_modes(args):157 """158 Test different optimization modes with verification.159 """160 parser = argparse.ArgumentParser(description="Test PyTorch Optimization Modes")161 parser.add_argument("--model_type", type=str, required=True, help="Model type")162 parser.add_argument("--config_path", type=str, required=True, help="Config path")163 parser.add_argument("--start_check_point", type=str, required=True, help="Checkpoint path (.ckpt)")164 parser.add_argument("--device", type=str, default='cuda:0', help="Device")165 166 if args is None:167 args = parser.parse_args()168 else:169 args = parser.parse_args(args)170 171 print("="*60)172 print("PyTorch Optimization Mode Test")173 print("="*60)174 175 # Load model176 print("\n📦 Loading model...")177 model, config = get_model_from_config(args.model_type, args.config_path)178 model = load_checkpoint(args.start_check_point, model, args.device)179 180 chunk_size = config.audio.chunk_size181 input_shape = (1, 2, chunk_size)182 dummy_input = torch.randn(*input_shape).to(args.device)183 184 # Test each optimization mode185 modes = ['default', 'compile', 'channels_last']186 outputs = {}187 188 for mode in modes:189 print(f"\n{'='*60}")190 print(f"Testing: {mode}")191 print('='*60)192 193 try:194 backend = PyTorchBackend(device=args.device, optimize_mode=mode)195 196 if mode == 'jit':197 backend.optimize_model(model, example_input=dummy_input, use_amp=True)198 else:199 backend.optimize_model(200 model, 201 use_amp=True,202 use_channels_last=(mode == 'channels_last')203 )204 205 # Run inference206 with torch.no_grad():207 output = backend(dummy_input)208 209 outputs[mode] = output210 print(f"✓ {mode} successful")211 print(f" Output shape: {output.shape}")212 print(f" Output range: [{output.min().item():.6f}, {output.max().item():.6f}]")213 214 except Exception as e:215 print(f"✗ {mode} failed: {e}")216 outputs[mode] = None217 218 # Verify outputs match219 print("\n" + "="*60)220 print("🔍 Output Verification")221 print("="*60)222 223 baseline_key = 'default'224 if baseline_key in outputs and outputs[baseline_key] is not None:225 baseline_output = outputs[baseline_key]226 227 for mode, output in outputs.items():228 if mode != baseline_key and output is not None:229 diff = torch.abs(baseline_output - output)230 max_diff = torch.max(diff).item()231 mean_diff = torch.mean(diff).item()232 233 print(f"\n{mode} vs {baseline_key}:")234 print(f" Max difference: {max_diff:.6f}")235 print(f" Mean difference: {mean_diff:.6f}")236 237 if max_diff < 1e-3:238 print(f" ✓ Outputs match within tolerance")239 else:240 print(f" ⚠ Warning: Large difference detected!")241 242 print("\n✅ Test completed!")243 244 245if __name__ == "__main__":246 import sys247 248 if len(sys.argv) > 1 and sys.argv[1] == 'test':249 sys.argv.pop(1)250 test_optimization_modes(None)251 else:252 benchmark_optimization_modes(None)253 