#!/usr/bin/env python """ Comprehensive import verification script for Suno GPT environment Tests all critical packages and their functionality """ import sys import importlib import traceback from typing import List, Tuple def test_basic_imports() -> List[Tuple[str, str, bool, str]]: """Test basic package imports""" packages = [ ("torch", "PyTorch"), ("torchvision", "TorchVision"), ("torchaudio", "TorchAudio"), ("transformers", "Transformers"), ("wandb", "Weights & Biases"), ("deepspeed", "DeepSpeed"), ("pytorch_lightning", "PyTorch Lightning"), ("nnAudio", "nnAudio"), ("auraloss", "Auraloss"), ("encodec", "Encodec"), ("flash_attn", "Flash Attention"), ("suno_utils", "Suno Utils"), ("g2p_en", "G2P English"), ("phonemizer", "Phonemizer"), ("sentencepiece", "SentencePiece"), ("tiktoken", "TikToken"), ("better_profanity", "Better Profanity"), ("einops", "Einops"), ("torchsde", "Torch SDE"), ("modal", "Modal"), ("ninja", "Ninja"), ("scipy", "SciPy"), ("numpy", "NumPy"), ("pandas", "Pandas"), ("tqdm", "TQDM"), ("joblib", "Joblib"), ("ffmpeg", "FFmpeg-Python"), ] results = [] for package, name in packages: try: mod = importlib.import_module(package) version = getattr(mod, '__version__', 'unknown') results.append((name, version, True, "OK")) except ImportError as e: results.append((name, "N/A", False, str(e)[:50])) except Exception as e: results.append((name, "N/A", False, f"Error: {str(e)[:50]}")) return results def test_cuda_functionality(): """Test CUDA and GPU functionality""" print("\n" + "="*60) print("CUDA and GPU Tests") print("="*60) try: import torch print(f"PyTorch version: {torch.__version__}") print(f"CUDA available: {torch.cuda.is_available()}") if torch.cuda.is_available(): print(f"CUDA version: {torch.version.cuda}") print(f"cuDNN version: {torch.backends.cudnn.version()}") print(f"Number of GPUs: {torch.cuda.device_count()}") for i in range(torch.cuda.device_count()): props = torch.cuda.get_device_properties(i) print(f"GPU {i}: {props.name}") print(f" Memory: {props.total_memory / 1024**3:.1f} GB") print(f" Compute Capability: {props.major}.{props.minor}") # Test basic CUDA operation try: x = torch.randn(100, 100).cuda() y = torch.randn(100, 100).cuda() z = torch.matmul(x, y) print("✓ CUDA tensor operations working") except Exception as e: print(f"✗ CUDA tensor operations failed: {e}") else: print("⚠ CUDA not available - CPU only mode") except ImportError: print("✗ PyTorch not available") except Exception as e: print(f"✗ Error testing CUDA: {e}") def test_flash_attention(): """Test Flash Attention functionality""" print("\n" + "="*60) print("Flash Attention Test") print("="*60) try: import torch import flash_attn from flash_attn import flash_attn_func print(f"Flash Attention version: {flash_attn.__version__}") # Test basic flash attention operation if torch.cuda.is_available(): batch, heads, seq_len, dim = 2, 8, 128, 64 q = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16) k = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16) v = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16) try: out = flash_attn_func(q, k, v) print(f"✓ Flash Attention working - output shape: {out.shape}") except Exception as e: print(f"✗ Flash Attention operation failed: {e}") else: print("⚠ Skipping Flash Attention test - CUDA not available") except ImportError as e: print(f"✗ Flash Attention not available: {e}") except Exception as e: print(f"✗ Error testing Flash Attention: {e}") def test_suno_utils(): """Test suno_utils functionality""" print("\n" + "="*60) print("Suno Utils Test") print("="*60) try: import suno_utils print("✓ suno_utils imported successfully") # Try to import some submodules test_modules = [ "suno_utils.audio", "suno_utils.io", "suno_utils.modeling", ] for module in test_modules: try: importlib.import_module(module) print(f"✓ {module} available") except ImportError: print(f"⚠ {module} not available") except ImportError as e: print(f"✗ suno_utils not available: {e}") except Exception as e: print(f"✗ Error testing suno_utils: {e}") def test_audio_packages(): """Test audio processing packages""" print("\n" + "="*60) print("Audio Package Tests") print("="*60) # Test nnAudio try: import nnAudio import torch from nnAudio.features import STFT if torch.cuda.is_available(): stft = STFT(n_fft=2048, hop_length=512).cuda() test_audio = torch.randn(1, 16000).cuda() spec = stft(test_audio) print(f"✓ nnAudio STFT working - output shape: {spec.shape}") else: print("⚠ nnAudio test skipped - CUDA not available") except Exception as e: print(f"⚠ nnAudio test failed: {e}") # Test auraloss try: import auraloss import torch if torch.cuda.is_available(): loss_fn = auraloss.freq.MultiResolutionSTFTLoss() x = torch.randn(1, 1, 16000).cuda() y = torch.randn(1, 1, 16000).cuda() loss = loss_fn(x, y) print(f"✓ Auraloss working - loss value: {loss.item():.4f}") else: print("⚠ Auraloss test skipped - CUDA not available") except Exception as e: print(f"⚠ Auraloss test failed: {e}") def main(): """Main verification function""" print("="*60) print("Suno GPT Environment Verification") print("="*60) print(f"Python: {sys.version}") print(f"Executable: {sys.executable}") print("="*60) # Test basic imports print("\nPackage Import Tests:") print("-"*40) results = test_basic_imports() # Print results in a nice table max_name_len = max(len(name) for name, _, _, _ in results) success_count = 0 failed_packages = [] for name, version, success, message in results: status = "✓" if success else "✗" if success: success_count += 1 print(f"{status} {name:<{max_name_len}} {version:<15}") else: failed_packages.append(name) print(f"{status} {name:<{max_name_len}} {'FAILED':<15} {message}") print(f"\nSummary: {success_count}/{len(results)} packages imported successfully") if failed_packages: print(f"\n⚠ Failed packages: {', '.join(failed_packages)}") # Run functionality tests test_cuda_functionality() test_flash_attention() test_suno_utils() test_audio_packages() # Final summary print("\n" + "="*60) if not failed_packages: print("✓ All critical packages are available!") return 0 else: print(f"⚠ Some packages failed to import. Please check the logs.") return 1 if __name__ == "__main__": sys.exit(main())