# setup.py

import re
import subprocess

from pathlib import Path
from typing import Literal

from setuptools import find_packages
from setuptools import setup


def get_project_version() -> str:
    """Read the public package version shared with liger-cute-kernels."""
    pyproject = Path(__file__).resolve().parent / "pyproject.toml"
    match = re.search(r'(?m)^version\s*=\s*"([^"]+)"\s*$', pyproject.read_text())
    if match is None:
        raise RuntimeError(f"could not read the project version from {pyproject}")
    return match.group(1)


def get_default_dependencies():
    """Determine the appropriate dependencies based on detected hardware."""
    platform = get_platform()

    if platform in ["cuda", "cpu"]:
        return [
            "torch>=2.1.2",
            "triton>=2.3.1",
        ]
    elif platform == "rocm":
        return [
            "triton>=3.0.0",
        ]
    elif platform == "xpu":
        return [
            "torch>=2.6.0",
        ]
    elif platform == "npu":
        return ["torch==2.9.0", "torch_npu==2.9.0", "triton-ascend==3.2.2"]


def get_optional_dependencies():
    """Get optional dependency groups."""
    # cuTile kernels use CompilerOptions.num_worker_warps (replace_hints / @ct.kernel),
    # which only exists in cuda-tile >= 1.4.0. Pin the floor to 1.5.0 (validated) so the
    # resolver can't backtrack to an older cuda-tile whose CompilerOptions lacks that
    # field (which raises "unexpected keyword argument 'num_worker_warps'" at import).
    cutile_deps = [
        "cuda-tile>=1.5.0",
    ]
    cutile_tileiras_deps = [
        "cuda-tile[tileiras]>=1.5.0",
    ]
    cutedsl_deps = [
        "nvidia-cutlass-dsl>=4.6.0",
        # Lets compiled CuTe DSL kernels take PyTorch tensors directly instead of
        # marshalling each one through DLPack per call. The kernels fall back to
        # the marshalling launch when it is absent, but on short kernels that
        # per-call cost dominates: RMSNorm forward measured 53us -> 15us on B200.
        "apache-tvm-ffi>=0.1.0",
    ]
    lck_deps = [
        f"liger-cute-kernels=={get_project_version()}",
    ]
    dev_deps = [
        "transformers>=4.52.0",
        "matplotlib>=3.7.2",
        "ruff>=0.12.0,<0.16.0",
        "pytest>=7.1.2",
        "pytest-xdist",
        "pytest-cov",
        "pytest-asyncio",
        "pytest-rerunfailures",
        "datasets>=2.19.2",
        "seaborn",
        "mkdocs-material",
        "torchvision>=0.20",
        "prek>=0.2.28",
    ]
    return {
        "cutile": cutile_deps,
        "cutile-tileiras": cutile_tileiras_deps,
        "cutedsl": cutedsl_deps,
        "lck": lck_deps,
        "dev": dev_deps,
    }


def is_xpu_available():
    """
    Check if Intel XPU is available.
    xpu-smi is often missing right now.
    """
    try:
        subprocess.run(["xpu-smi"], check=True)
        return True
    except (subprocess.SubprocessError, FileNotFoundError):
        pass

    try:
        result = subprocess.run("sycl-ls", check=True, capture_output=True, shell=True)
        if "level_zero:gpu" in result.stdout.decode():
            return True
    except (subprocess.SubprocessError, FileNotFoundError):
        pass

    return False


def is_ascend_available() -> bool:
    """Best-effort Ascend detection.

    Checks for common Ascend environment variables and a possible `npu-smi`
    utility if present.
    """
    try:
        subprocess.run(["npu-smi", "info"], check=True)
        return True
    except (subprocess.SubprocessError, FileNotFoundError):
        pass
    return False


def get_platform() -> Literal["cuda", "rocm", "cpu", "xpu", "npu"]:
    """
    Detect whether the system has NVIDIA or AMD GPU without torch dependency.
    """
    # Try nvidia-smi first
    try:
        subprocess.run(["nvidia-smi"], check=True)
        print("NVIDIA GPU detected")
        return "cuda"
    except (subprocess.SubprocessError, FileNotFoundError):
        # If nvidia-smi fails, check for ROCm
        try:
            subprocess.run(["rocm-smi"], check=True)
            print("ROCm GPU detected")
            return "rocm"
        except (subprocess.SubprocessError, FileNotFoundError):
            if is_xpu_available():
                print("Intel GPU detected")
                return "xpu"
            elif is_ascend_available():
                print("Ascend NPU detected")
                return "npu"
            else:
                print("No GPU detected")
                return "cpu"


setup(
    name="liger_kernel",
    package_dir={"": "src"},
    packages=find_packages(where="src"),
    install_requires=get_default_dependencies(),
    extras_require=get_optional_dependencies(),
    classifiers=[
        "Development Status :: 5 - Production/Stable",
        "Intended Audience :: Developers",
        "Intended Audience :: Education",
        "Intended Audience :: Science/Research",
        "Programming Language :: Python :: 3",
        "Topic :: Scientific/Engineering :: Artificial Intelligence",
        "Topic :: Software Development :: Libraries :: Python Modules",
        "License :: OSI Approved :: BSD-2-Clause Software License",
        "Operating System :: OS Independent",
    ],
)
