| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import importlib |
| import sys |
| from argparse import ArgumentParser |
|
|
| parser = ArgumentParser() |
| parser.add_argument("--training", action="store_true", help="Check training packages") |
| args = parser.parse_args() |
|
|
|
|
| def _test_flash_attn(): |
| import torch |
|
|
| |
| from flash_attn import flash_attn_func |
|
|
| device = "cuda" |
| dtype = torch.float16 |
|
|
| batch_size = 2 |
| seqlen = 512 |
| num_heads = 8 |
| head_dim = 64 |
|
|
| q = torch.randn(batch_size, seqlen, num_heads, head_dim, device=device, dtype=dtype, requires_grad=True) |
| k = torch.randn(batch_size, seqlen, num_heads, head_dim, device=device, dtype=dtype, requires_grad=True) |
| v = torch.randn(batch_size, seqlen, num_heads, head_dim, device=device, dtype=dtype, requires_grad=True) |
|
|
| out = flash_attn_func( |
| q, |
| k, |
| v, |
| dropout_p=0.0, |
| softmax_scale=1.0 / (head_dim**0.5), |
| causal=True, |
| window_size=(-1, -1), |
| softcap=0.0, |
| ) |
|
|
|
|
| def _flash_attn_is_ok(): |
| try: |
| _test_flash_attn() |
| except ImportError: |
| return False |
| except RuntimeError: |
| return False |
| return True |
|
|
|
|
| def check_packages(package_list, success_status=True): |
| def print_success(package, version=None): |
| if version: |
| print(f"\033[92m[SUCCESS]\033[0m {package} found (v{version})") |
| else: |
| print(f"\033[92m[SUCCESS]\033[0m {package} found") |
|
|
| def print_error(message): |
| print(f"\033[91m[ERROR]\033[0m {message}") |
|
|
| for package in package_list: |
| if isinstance(package, tuple): |
| found = False |
| for alt_package in package: |
| try: |
| module = importlib.import_module(alt_package) |
| version = getattr(module, "__version__", None) |
| print_success(alt_package, version) |
| found = True |
| break |
| except ImportError: |
| continue |
| if not found: |
| print_error(f"None of the alternative packages found: \033[93m{', '.join(package)}\033[0m") |
| success_status = False |
| elif package == "apex": |
| try: |
| module = importlib.import_module(package) |
| version = getattr(module, "__version__", None) |
| print_success(package, version) |
| try: |
| from apex import multi_tensor_apply |
|
|
| print_success("apex.multi_tensor_apply") |
| except ImportError: |
| print_error("apex.multi_tensor_apply not found") |
| success_status = False |
| except ImportError: |
| print_error("apex not found") |
| success_status = False |
| elif package == "transformer_engine": |
| try: |
| module = importlib.import_module(package) |
| version = getattr(module, "__version__", None) |
| print_success(package, version) |
| try: |
| import transformer_engine.pytorch |
|
|
| print_success("transformer_engine.pytorch") |
| except ImportError: |
| print_error("transformer_engine.pytorch not found") |
| success_status = False |
| except ImportError: |
| print_error("transformer_engine not found") |
| success_status = False |
| else: |
| try: |
| module = importlib.import_module(package) |
| version = getattr(module, "__version__", None) |
| print_success(package, version) |
| except ImportError as e: |
| print_error(f"Package not successfully imported: \033[93m{package}\033[0m") |
| success_status = False |
|
|
| if _flash_attn_is_ok(): |
| print(f"\033[92m[SUCCESS]\033[0m flash_attn_func succeeds") |
| else: |
| print(f"\033[91m[ERROR]\033[0m flash_attn_func fails") |
| success_status = False |
|
|
| return success_status |
|
|
|
|
| if not (sys.version_info.major == 3 and sys.version_info.minor >= 10): |
| detected = f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}" |
| print(f"\033[91m[ERROR]\033[0m Python 3.10+ is required. You have: \033[93m{detected}\033[0m") |
| sys.exit(1) |
|
|
| print("Attempting to import critical packages...") |
|
|
| packages = [ |
| "torch", |
| "torchvision", |
| "diffusers", |
| "transformers", |
| "transformer_engine", |
| "megatron.core", |
| ("flash_attn", "flash_attn_interface"), |
| "natten", |
| ] |
| packages_training = [ |
| "apex", |
| ] |
|
|
| all_success = check_packages(packages) |
| if args.training: |
| training_success = check_packages(packages_training) |
| if not training_success: |
| print("\033[93m[WARNING]\033[0m Training packages not found. Training features will be unavailable.") |
|
|
| if all_success: |
| print("-----------------------------------------------------------") |
| print("\033[92m[SUCCESS]\033[0m Cosmos-predict2 environment setup is successful!") |
|
|