-
Notifications
You must be signed in to change notification settings - Fork 298
Expand file tree
/
Copy pathpyproject.toml
More file actions
142 lines (134 loc) · 4.96 KB
/
Copy pathpyproject.toml
File metadata and controls
142 lines (134 loc) · 4.96 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
[project]
name = "nvidia-cudnn-frontend"
dynamic = ["version"]
description = "NVIDIA cuDNN Frontend — Python and C++ Graph API with SOTA attention (SDPA / Flash Attention), MoE grouped GEMM fusions, and FP8/MXFP8 kernels for Hopper, Blackwell, and Rubin GPUs."
readme = "README.md"
requires-python = ">=3.10"
license = {text = "Apache-2.0 AND MIT"}
keywords = [
"cudnn",
"cuda",
"gpu",
"nvidia",
"deep-learning",
"attention",
"sdpa",
"flash-attention",
"transformer",
"moe",
"mixture-of-experts",
"grouped-gemm",
"fp8",
"mxfp8",
"blackwell",
"rubin",
"hopper",
"pytorch",
"kernel",
"graph-api",
]
classifiers = [
"Development Status :: 5 - Production/Stable",
"Intended Audience :: Developers",
"Intended Audience :: Science/Research",
"License :: OSI Approved :: Apache Software License",
"License :: OSI Approved :: MIT License",
"Operating System :: POSIX :: Linux",
"Operating System :: Microsoft :: Windows",
"Programming Language :: C++",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Programming Language :: Python :: 3.14",
"Environment :: GPU :: NVIDIA CUDA",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"Topic :: Software Development :: Libraries :: Python Modules",
]
dependencies = [
"nvidia-cutlass-dsl[cu13]>=4.6.2",
"apache-tvm-ffi>=0.1.11",
]
[project.urls]
"Homepage" = "https://github.com/NVIDIA/cudnn-frontend"
"Documentation" = "https://docs.nvidia.com/deeplearning/cudnn/frontend/latest/"
"Blog" = "https://nvidia.github.io/cudnn-frontend/"
"Repository" = "https://github.com/NVIDIA/cudnn-frontend"
"Bug Tracker" = "https://github.com/NVIDIA/cudnn-frontend/issues"
"Release Notes" = "https://github.com/NVIDIA/cudnn-frontend/releases"
[project.optional-dependencies]
cutedsl = [
# nvidia-cutlass-dsl and apache-tvm-ffi moved to the required [project]
# dependencies above; this extra is what is left of the old CuTeDSL group.
# Naming cuda-python here is belt-and-braces rather than load-bearing: the
# required nvidia-cutlass-dsl pulls nvidia-cutlass-dsl-libs-base, which itself
# requires cuda-python>=12.8, so a plain `pip install nvidia-cudnn-frontend`
# already has it. Keeping it spelled out pins the dependency to this package
# instead of to a transitive detail of the DSL's packaging, and keeps
# `pip install nvidia-cudnn-frontend[cutedsl]` resolving for everyone who
# still writes it that way.
"cuda-python",
]
comm = [
# Communication runtimes are grouped by backend so future distributed
# operation graphs can reuse the same installation extra.
"nvshmem4py-cu13>=0.3.1",
]
cutile = [
# The cuTile linear-attention engines. Base cuda-tile only -- its [tileiras]
# extra pins cuda-toolkit>=13.2,<13.4, and that upper bound would cap the
# whole environment's toolkit and shut out CUDA 12 entirely. Without it,
# cuda.tile falls back to a system `tileiras`, the same way this package
# already leaves GPU wheels to the user. Engines that cannot import the
# runtime decline in check_support, so a missing compiler costs those
# engines and nothing else.
"cuda-tile>=1.4",
]
triton = [
# Framework-specific Triton FE OSS kernels. Torch remains in the explicit
# torch dependency group so callers can choose the matching CUDA wheel.
"triton>=3.7.0; sys_platform == 'linux' and python_version >= '3.10'",
]
[dependency-groups]
dev = [
"jupyter",
"numpy",
"pybind11[global]>=2.13,<3",
"pytest",
"pytest-xdist",
"cuda-python",
"looseversion",
"black==26.3.1",
"clang-format==21.1.6",
]
# Per-framework groups for the type-erased CuTeDSL GEMM APIs (JAX is supported
# by the dense fusions: amax, swiglu, srelu, dsrelu): install with
# `pip install --group torch` / `pip install --group jax`. GPU wheels
# (torch cuXX, jax[cuda12]/[cuda13]) are left to the user.
torch = [
"torch",
"torch-c-dlpack-ext",
]
# The jax.jit-compatible XLA custom-call entry points (cudnn.jax, e.g.
# gemm_amax_jax_sm100) build on the CuTeDSL JAX extensions (cutlass.jax,
# shipped with nvidia-cutlass-dsl), which require jax >= 0.5.
jax = [
"jax>=0.5",
]
[build-system]
requires = ["setuptools>=64", "cmake>=3.18", "ninja==1.11.1.1", "pybind11[global]>=2.13,<3"]
build-backend = "setuptools.build_meta"
[tool.setuptools]
packages = {find = {where = ["python", "."], include = ["cudnn*", "include"], namespaces = true}}
package-dir = {"" = "python", "include" = "include"}
include-package-data = true
[tool.setuptools.dynamic]
version = {attr = "cudnn.__version__"}
[tool.setuptools.package-data]
include = ["**/*"]
"cudnn.moe_ep._megamoe_backend.cutedsl_src" = [
"LICENSE.Apache-2.0",
"VENDOR.md",
]
"cudnn.linear_attention.cake.kernels" = ["*.cu", "SHA256SUMS", "UPSTREAM.md"]