CoolFace
Apppublic

ASesYusuf1/SESA_Audio_Separation

sourceHugging Facemitupdated 6mo agoView on Hugging Face
14likes
benchmark_pytorch.py253 linesDownload Raw Back to root
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