diff --git a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.10.____cpython.yaml b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.10.____cpython.yaml index 4e19191..072ad5d 100644 --- a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.10.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.10.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.11.____cpython.yaml b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.11.____cpython.yaml index f77b690..f530a5b 100644 --- a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.11.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.11.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.12.____cpython.yaml b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.12.____cpython.yaml index de20559..866f345 100644 --- a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.12.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.12.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.9.____cpython.yaml b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.9.____cpython.yaml index 69b099e..df5a6d3 100644 --- a/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.9.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version11cuda_compilernvcccuda_compiler_version11.8cxx_compiler_version11python3.9.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.10.____cpython.yaml b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.10.____cpython.yaml index 7b46f9f..9562803 100644 --- a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.10.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.10.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.11.____cpython.yaml b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.11.____cpython.yaml index 3ceb8f2..f53693a 100644 --- a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.11.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.11.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.12.____cpython.yaml b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.12.____cpython.yaml index a63ae35..f45ddea 100644 --- a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.12.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.12.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.9.____cpython.yaml b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.9.____cpython.yaml index 06a95ad..de13864 100644 --- a/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.9.____cpython.yaml +++ b/.ci_support/linux_64_c_compiler_version12cuda_compilercuda-nvcccuda_compiler_version12.0cxx_compiler_version12python3.9.____cpython.yaml @@ -7,7 +7,7 @@ c_stdlib: c_stdlib_version: - '2.17' cdt_name: -- cos7 +- conda channel_sources: - conda-forge channel_targets: @@ -40,7 +40,6 @@ zip_keys: - - c_compiler_version - cxx_compiler_version - c_stdlib_version - - cdt_name - cuda_compiler - cuda_compiler_version - docker_image diff --git a/.ci_support/migrations/python312.yaml b/.ci_support/migrations/python312.yaml deleted file mode 100644 index 784a0a2..0000000 --- a/.ci_support/migrations/python312.yaml +++ /dev/null @@ -1,38 +0,0 @@ -migrator_ts: 1695046563 -__migrator: - migration_number: 1 - operation: key_add - primary_key: python - ordering: - python: - - 3.6.* *_cpython - - 3.7.* *_cpython - - 3.8.* *_cpython - - 3.9.* *_cpython - - 3.10.* *_cpython - - 3.11.* *_cpython - - 3.12.* *_cpython # new entry - - 3.6.* *_73_pypy - - 3.7.* *_73_pypy - - 3.8.* *_73_pypy - - 3.9.* *_73_pypy - paused: false - longterm: True - pr_limit: 30 - max_solver_attempts: 6 # this will make the bot retry "not solvable" stuff 6 times - exclude: - # this shouldn't attempt to modify the python feedstocks - - python - - pypy3.6 - - pypy-meta - - cross-python - - python_abi - exclude_pinned_pkgs: false - -python: - - 3.12.* *_cpython -# additional entries to add for zip_keys -numpy: - - 1.26 -python_impl: - - cpython diff --git a/.github/workflows/automerge.yml b/.github/workflows/automerge.yml deleted file mode 100644 index 0535f6a..0000000 --- a/.github/workflows/automerge.yml +++ /dev/null @@ -1,17 +0,0 @@ -on: - status: {} - check_suite: - types: - - completed - -jobs: - automerge-action: - runs-on: ubuntu-latest - name: automerge - steps: - - name: automerge-action - id: automerge-action - uses: conda-forge/automerge-action@main - with: - github_token: ${{ secrets.GITHUB_TOKEN }} - rerendering_github_token: ${{ secrets.RERENDERING_GITHUB_TOKEN }} diff --git a/.github/workflows/conda-build.yml b/.github/workflows/conda-build.yml index 1623e44..ecff2ce 100644 --- a/.github/workflows/conda-build.yml +++ b/.github/workflows/conda-build.yml @@ -16,7 +16,7 @@ jobs: build: name: ${{ matrix.CONFIG }} runs-on: ${{ matrix.runs_on }} - timeout-minutes: 540 + timeout-minutes: 1080 strategy: fail-fast: false matrix: @@ -117,12 +117,6 @@ jobs: fi ./.scripts/run_osx_build.sh - - name: Install Miniconda for windows - uses: conda-incubator/setup-miniconda@a4260408e20b96e80095f42ff7f1a15b27dd94ca # v3.0.4 - with: - miniforge-version: latest - if: matrix.os == 'windows' - - name: Build on windows shell: cmd run: | @@ -131,6 +125,7 @@ jobs: set "sha=%GITHUB_SHA%" call ".scripts\run_win_build.bat" env: + MINIFORGE_HOME: D:\Miniforge PYTHONUNBUFFERED: 1 CONFIG: ${{ matrix.CONFIG }} CI: github_actions diff --git a/.github/workflows/webservices.yml b/.github/workflows/webservices.yml deleted file mode 100644 index d6f06b5..0000000 --- a/.github/workflows/webservices.yml +++ /dev/null @@ -1,13 +0,0 @@ -on: repository_dispatch - -jobs: - webservices: - runs-on: ubuntu-latest - name: webservices - steps: - - name: webservices - id: webservices - uses: conda-forge/webservices-dispatch-action@main - with: - github_token: ${{ secrets.GITHUB_TOKEN }} - rerendering_github_token: ${{ secrets.RERENDERING_GITHUB_TOKEN }} diff --git a/.scripts/build_steps.sh b/.scripts/build_steps.sh index 9123720..f8051ab 100755 --- a/.scripts/build_steps.sh +++ b/.scripts/build_steps.sh @@ -31,13 +31,13 @@ pkgs_dirs: solver: libmamba CONDARC +mv /opt/conda/conda-meta/history /opt/conda/conda-meta/history.$(date +%Y-%m-%d-%H-%M-%S) +echo > /opt/conda/conda-meta/history +micromamba install --root-prefix ~/.conda --prefix /opt/conda \ + --yes --override-channels --channel conda-forge --strict-channel-priority \ + pip python=3.12 conda-build conda-forge-ci-setup=4 "conda-build>=24.1" export CONDA_LIBMAMBA_SOLVER_NO_CHANNELS_FROM_INSTALLED=1 -mamba install --update-specs --yes --quiet --channel conda-forge --strict-channel-priority \ - pip mamba conda-build conda-forge-ci-setup=4 "conda-build>=24.1" -mamba update --update-specs --yes --quiet --channel conda-forge --strict-channel-priority \ - pip mamba conda-build conda-forge-ci-setup=4 "conda-build>=24.1" - # set up the condarc setup_conda_rc "${FEEDSTOCK_ROOT}" "${RECIPE_ROOT}" "${CONFIG_FILE}" diff --git a/README.md b/README.md index 42d2d8b..22a9d81 100644 --- a/README.md +++ b/README.md @@ -22,6 +22,8 @@ Current release info | Name | Downloads | Version | Platforms | | --- | --- | --- | --- | | [![Conda Recipe](https://img.shields.io/badge/recipe-flash--attn-green.svg)](https://anaconda.org/conda-forge/flash-attn) | [![Conda Downloads](https://img.shields.io/conda/dn/conda-forge/flash-attn.svg)](https://anaconda.org/conda-forge/flash-attn) | [![Conda Version](https://img.shields.io/conda/vn/conda-forge/flash-attn.svg)](https://anaconda.org/conda-forge/flash-attn) | [![Conda Platforms](https://img.shields.io/conda/pn/conda-forge/flash-attn.svg)](https://anaconda.org/conda-forge/flash-attn) | +| [![Conda Recipe](https://img.shields.io/badge/recipe-flash--attn--fused--dense-green.svg)](https://anaconda.org/conda-forge/flash-attn-fused-dense) | [![Conda Downloads](https://img.shields.io/conda/dn/conda-forge/flash-attn-fused-dense.svg)](https://anaconda.org/conda-forge/flash-attn-fused-dense) | [![Conda Version](https://img.shields.io/conda/vn/conda-forge/flash-attn-fused-dense.svg)](https://anaconda.org/conda-forge/flash-attn-fused-dense) | [![Conda Platforms](https://img.shields.io/conda/pn/conda-forge/flash-attn-fused-dense.svg)](https://anaconda.org/conda-forge/flash-attn-fused-dense) | +| [![Conda Recipe](https://img.shields.io/badge/recipe-flash--attn--layer--norm-green.svg)](https://anaconda.org/conda-forge/flash-attn-layer-norm) | [![Conda Downloads](https://img.shields.io/conda/dn/conda-forge/flash-attn-layer-norm.svg)](https://anaconda.org/conda-forge/flash-attn-layer-norm) | [![Conda Version](https://img.shields.io/conda/vn/conda-forge/flash-attn-layer-norm.svg)](https://anaconda.org/conda-forge/flash-attn-layer-norm) | [![Conda Platforms](https://img.shields.io/conda/pn/conda-forge/flash-attn-layer-norm.svg)](https://anaconda.org/conda-forge/flash-attn-layer-norm) | Installing flash-attn ===================== @@ -33,16 +35,16 @@ conda config --add channels conda-forge conda config --set channel_priority strict ``` -Once the `conda-forge` channel has been enabled, `flash-attn` can be installed with `conda`: +Once the `conda-forge` channel has been enabled, `flash-attn, flash-attn-fused-dense, flash-attn-layer-norm` can be installed with `conda`: ``` -conda install flash-attn +conda install flash-attn flash-attn-fused-dense flash-attn-layer-norm ``` or with `mamba`: ``` -mamba install flash-attn +mamba install flash-attn flash-attn-fused-dense flash-attn-layer-norm ``` It is possible to list all of the versions of `flash-attn` available on your platform with `conda`: diff --git a/conda-forge.yml b/conda-forge.yml index 5dbd5ef..da58b5b 100644 --- a/conda-forge.yml +++ b/conda-forge.yml @@ -5,7 +5,7 @@ conda_build: error_overlinking: true conda_forge_output_validation: true github_actions: - timeout_minutes: 540 + timeout_minutes: 1080 self_hosted: true triggers: - push diff --git a/recipe/meta.yaml b/recipe/meta.yaml index 461ef3b..b7a57d5 100644 --- a/recipe/meta.yaml +++ b/recipe/meta.yaml @@ -1,19 +1,18 @@ -{% set name = "flash-attn" %} {% set version = "2.6.3" %} package: - name: {{ name|lower }} + name: flash-attn-split version: {{ version }} source: - - url: https://pypi.io/packages/source/{{ name[0] }}/{{ name }}/flash_attn-{{ version }}.tar.gz + - url: https://pypi.org/packages/source/f/flash-attn/flash_attn-{{ version }}.tar.gz sha256: 5bfae9500ad8e7d2937ebccb4906f3bc464d1bf66eedd0e4adabd520811c7b52 # Overwrite with a simpler build script that doesn't try to revend pre-compiled binaries - path: pyproject.toml - path: setup.py build: - number: 1 + number: 2 script: {{ PYTHON }} -m pip install . -vvv --no-deps --no-build-isolation script_env: # Limit MAX_JOBS in order to prevent runners from crashing @@ -23,12 +22,8 @@ build: skip: true # [not linux] skip: true # [py==313] # Skip until pytorch dependency on setuptools is fixed # debugging skips below - # skip: true # [py!=313] - # skip: true # [cuda_compiler_version != "12.0"] - ignore_run_exports_from: - - libcublas-dev # [(cuda_compiler_version or "").startswith("12")] - - libcusolver-dev # [(cuda_compiler_version or "").startswith("12")] - - libcusparse-dev # [(cuda_compiler_version or "").startswith("12")] + # skip: true # [py!=312] + # skip: true # [cuda_compiler_version != "11.8"] requirements: build: @@ -41,6 +36,7 @@ requirements: - cuda-version {{ cuda_compiler_version }} # same cuda for host and build - cuda-cudart-dev # [(cuda_compiler_version or "").startswith("12")] - libcublas-dev # [(cuda_compiler_version or "").startswith("12")] + - libcurand-dev # [(cuda_compiler_version or "").startswith("12")] - libcusolver-dev # [(cuda_compiler_version or "").startswith("12")] - libcusparse-dev # [(cuda_compiler_version or "").startswith("12")] - libtorch # required until pytorch run_exports libtorch @@ -49,18 +45,97 @@ requirements: - pytorch - pytorch =*=cuda* - setuptools - run: - - einops - - python - - pytorch =*=cuda* -test: - imports: - - flash_attn - commands: - - pip check - requires: - - pip +outputs: + - name: flash-attn + requirements: + build: + - {{ compiler('c') }} + - {{ compiler('cxx') }} + - {{ compiler('cuda') }} + - {{ stdlib('c') }} + host: + - cuda-version {{ cuda_compiler_version }} # same cuda for host and build + - cuda-cudart-dev # [(cuda_compiler_version or "").startswith("12")] + - python + - libtorch # required until pytorch run_exports libtorch + - pytorch + - pytorch =*=cuda* + run: + - einops + - python + - pytorch =*=cuda* + + files: + include: + - 'lib/python*/site-packages/flash_attn/**' + - 'lib/python*/site-packages/flash_attn-{{ version }}.dist-info/**' + - 'lib/python*/site-packages/flash_attn_2_cuda.cpython-*.so' + + test: + imports: + - flash_attn + commands: + - pip check + requires: + - pip + + - name: flash-attn-fused-dense + requirements: + build: + - {{ compiler('cxx') }} # needed for DSO checker + - {{ stdlib('c') }} # needed for DSO checker + - {{ compiler('cuda') }} + host: + - cuda-version {{ cuda_compiler_version }} # same cuda for host and build + - cuda-cudart-dev # [(cuda_compiler_version or "").startswith("12")] + - libcublas-dev # [(cuda_compiler_version or "").startswith("12")] + - python + - pytorch + - pytorch =*=cuda* + run: + - python + - {{ pin_subpackage('flash-attn', exact=True) }} + + files: + include: + - 'lib/python*/site-packages/fused_dense_lib.cpython-*.so' + + test: + imports: + - flash_attn.ops.fused_dense + commands: + - pip check + requires: + - pip + + - name: flash-attn-layer-norm + requirements: + build: + - {{ compiler('cxx') }} # needed for DSO checker + - {{ stdlib('c') }} # needed for DSO checker + - {{ compiler('cuda') }} + host: + - cuda-version {{ cuda_compiler_version }} # same cuda for host and build + - cuda-cudart-dev # [(cuda_compiler_version or "").startswith("12")] + - python + - pytorch + - pytorch =*=cuda* + run: + - python + - {{ pin_subpackage('flash-attn', exact=True) }} + + files: + include: + - 'lib/python*/site-packages/dropout_layer_norm.cpython-*.so' + + test: + imports: + - flash_attn.ops.layer_norm + commands: + - pip check + requires: + - pip about: home: https://github.com/Dao-AILab/flash-attention @@ -71,6 +146,7 @@ about: - LICENSE_CUTLASS.txt extra: + feedstock-name: flash-attn recipe-maintainers: - carterbox - weiji14 diff --git a/recipe/setup.py b/recipe/setup.py index b152e80..56059a9 100644 --- a/recipe/setup.py +++ b/recipe/setup.py @@ -10,6 +10,7 @@ """ import pathlib +import sys from setuptools import setup, find_packages from torch.utils.cpp_extension import BuildExtension, CUDAExtension @@ -132,11 +133,121 @@ # "-DFLASHATTENTION_DISABLE_LOCAL", ], }, + extra_link_args=["-Wl,--strip-all", "-Wl,--no-undefined"], include_dirs=[ _this_dir / "csrc" / "flash_attn", _this_dir / "csrc" / "flash_attn" / "src", _this_dir / "csrc" / "cutlass" / "include", ], + libraries=[ + "cudart", + f"python{sys.version_info[0]}.{sys.version_info[1]}", + ], + ), + CUDAExtension( + name="fused_dense_lib", + sources=[ + "csrc/fused_dense_lib/fused_dense.cpp", + "csrc/fused_dense_lib/fused_dense_cuda.cu", + ], + extra_compile_args={ + "cxx": [ + "-O3", + ], + "nvcc": ["-O3"], + }, + extra_link_args=["-Wl,--strip-all", "-Wl,--no-undefined"], + libraries=[ + "cudart", + "cublas", + "cublasLt", + f"python{sys.version_info[0]}.{sys.version_info[1]}", + ], + ), + CUDAExtension( + name="dropout_layer_norm", + sources=[ + "csrc/layer_norm/ln_api.cpp", + "csrc/layer_norm/ln_fwd_256.cu", + "csrc/layer_norm/ln_bwd_256.cu", + "csrc/layer_norm/ln_fwd_512.cu", + "csrc/layer_norm/ln_bwd_512.cu", + "csrc/layer_norm/ln_fwd_768.cu", + "csrc/layer_norm/ln_bwd_768.cu", + "csrc/layer_norm/ln_fwd_1024.cu", + "csrc/layer_norm/ln_bwd_1024.cu", + "csrc/layer_norm/ln_fwd_1280.cu", + "csrc/layer_norm/ln_bwd_1280.cu", + "csrc/layer_norm/ln_fwd_1536.cu", + "csrc/layer_norm/ln_bwd_1536.cu", + "csrc/layer_norm/ln_fwd_2048.cu", + "csrc/layer_norm/ln_bwd_2048.cu", + "csrc/layer_norm/ln_fwd_2560.cu", + "csrc/layer_norm/ln_bwd_2560.cu", + "csrc/layer_norm/ln_fwd_3072.cu", + "csrc/layer_norm/ln_bwd_3072.cu", + "csrc/layer_norm/ln_fwd_4096.cu", + "csrc/layer_norm/ln_bwd_4096.cu", + "csrc/layer_norm/ln_fwd_5120.cu", + "csrc/layer_norm/ln_bwd_5120.cu", + "csrc/layer_norm/ln_fwd_6144.cu", + "csrc/layer_norm/ln_bwd_6144.cu", + "csrc/layer_norm/ln_fwd_7168.cu", + "csrc/layer_norm/ln_bwd_7168.cu", + "csrc/layer_norm/ln_fwd_8192.cu", + "csrc/layer_norm/ln_bwd_8192.cu", + "csrc/layer_norm/ln_parallel_fwd_256.cu", + "csrc/layer_norm/ln_parallel_bwd_256.cu", + "csrc/layer_norm/ln_parallel_fwd_512.cu", + "csrc/layer_norm/ln_parallel_bwd_512.cu", + "csrc/layer_norm/ln_parallel_fwd_768.cu", + "csrc/layer_norm/ln_parallel_bwd_768.cu", + "csrc/layer_norm/ln_parallel_fwd_1024.cu", + "csrc/layer_norm/ln_parallel_bwd_1024.cu", + "csrc/layer_norm/ln_parallel_fwd_1280.cu", + "csrc/layer_norm/ln_parallel_bwd_1280.cu", + "csrc/layer_norm/ln_parallel_fwd_1536.cu", + "csrc/layer_norm/ln_parallel_bwd_1536.cu", + "csrc/layer_norm/ln_parallel_fwd_2048.cu", + "csrc/layer_norm/ln_parallel_bwd_2048.cu", + "csrc/layer_norm/ln_parallel_fwd_2560.cu", + "csrc/layer_norm/ln_parallel_bwd_2560.cu", + "csrc/layer_norm/ln_parallel_fwd_3072.cu", + "csrc/layer_norm/ln_parallel_bwd_3072.cu", + "csrc/layer_norm/ln_parallel_fwd_4096.cu", + "csrc/layer_norm/ln_parallel_bwd_4096.cu", + "csrc/layer_norm/ln_parallel_fwd_5120.cu", + "csrc/layer_norm/ln_parallel_bwd_5120.cu", + "csrc/layer_norm/ln_parallel_fwd_6144.cu", + "csrc/layer_norm/ln_parallel_bwd_6144.cu", + "csrc/layer_norm/ln_parallel_fwd_7168.cu", + "csrc/layer_norm/ln_parallel_bwd_7168.cu", + "csrc/layer_norm/ln_parallel_fwd_8192.cu", + "csrc/layer_norm/ln_parallel_bwd_8192.cu", + ], + extra_compile_args={ + "cxx": ["-O3"], + "nvcc": [ + "-O3", + "-U__CUDA_NO_HALF_OPERATORS__", + "-U__CUDA_NO_HALF_CONVERSIONS__", + "-U__CUDA_NO_BFLOAT16_OPERATORS__", + "-U__CUDA_NO_BFLOAT16_CONVERSIONS__", + "-U__CUDA_NO_BFLOAT162_OPERATORS__", + "-U__CUDA_NO_BFLOAT162_CONVERSIONS__", + "--expt-relaxed-constexpr", + "--expt-extended-lambda", + "--use_fast_math", + ], + }, + extra_link_args=["-Wl,--strip-all", "-Wl,--no-undefined"], + include_dirs=[ + _this_dir / "csrc" / "layer_norm", + ], + libraries=[ + "cudart", + f"python{sys.version_info[0]}.{sys.version_info[1]}", + ], ), ], cmdclass={"build_ext": BuildExtension},