Repository navigation
build: update development PyTorch image to 26.09 #7725
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 49 commits
84f81ea
1ff951b
2560c5a
7640398
dac0cff
d58a15d
46d3cca
67ee88c
6ebe61d
e9f01a0
2f16cb1
71f91fc
64dd800
0e2c866
e1e0c85
87a35cd
a883731
85f8fc5
b028d4e
28f8046
e071ec3
74ad5bb
070712a
20838da
4906863
9b8300d
c9819a9
bd4fe43
ee07a6b
513efe4
45679b9
39bb2bc
718410b
0b07373
6fef589
16c1b9a
12b1c0a
3befeee
8f99759
4f4d214
3e3e659
88879c2
4546569
ec34db6
4114504
c5378c3
b73d6cb
25b5f63
64977df
d842b7a
090200d
fab15a7
a7a6534
23e74cc
896ea0e
46e7128
880404c
89fde7c
5b38ce0
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1 +1 @@ | ||
| nvcr.io/nvidia/pytorch:26.08-py3 | ||
| nvcr.io/nvidia/pytorch:26.09-py3 |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,47 @@ | ||
| diff --git a/flash_attn/cute/flash_bwd_sm100.py b/flash_attn/cute/flash_bwd_sm100.py | ||
| index bded7c05303..5cf6fb506c7 100644 | ||
| --- a/flash_attn/cute/flash_bwd_sm100.py | ||
| +++ b/flash_attn/cute/flash_bwd_sm100.py | ||
| @@ -14,7 +14,6 @@ | ||
| import cutlass.utils.blackwell_helpers as sm100_utils_basic | ||
| from cutlass.pipeline import PipelineAsync | ||
|
|
||
| -import quack.activation | ||
| from quack import layout_utils | ||
| from flash_attn.cute import utils | ||
| from flash_attn.cute.cute_dsl_utils import assume_tensor_aligned | ||
| @@ -3268,7 +3267,7 @@ def compute_loop( | ||
| utils.shuffle_sync(tSrdPsum, offset=2 * v + 1), | ||
| ) | ||
| tdPrdP_cur[2 * v], tdPrdP_cur[2 * v + 1] = ( | ||
| - quack.activation.sub_packed_f32x2( | ||
| + cute.arch.sub_packed_f32x2( | ||
| (tdPrdP_cur[2 * v], tdPrdP_cur[2 * v + 1]), dPsum_pair | ||
| ) | ||
| ) | ||
| diff --git a/flash_attn/cute/utils.py b/flash_attn/cute/utils.py | ||
| index 9feafab5f05..5391fa7cfed 100644 | ||
| --- a/flash_attn/cute/utils.py | ||
| +++ b/flash_attn/cute/utils.py | ||
| @@ -17,8 +17,6 @@ | ||
| from cutlass.cute.runtime import from_dlpack | ||
|
|
||
|
|
||
| -import quack.activation | ||
| - | ||
| _MIXER_ATTRS = ("__vec_size__",) | ||
|
|
||
|
|
||
| @@ -780,10 +778,10 @@ def ex2_emulation_2( | ||
| xy_rounded = cute.arch.add_packed_f32x2(xy_clamped, (fp32_round_int, fp32_round_int), rnd="rm") | ||
| # The integer floor of x & y are now in the last 8 bits of xy_rounded | ||
| # We want the next 2 ops to round to nearest even. The rounding mode is important. | ||
| - xy_rounded_back = quack.activation.sub_packed_f32x2( | ||
| + xy_rounded_back = cute.arch.sub_packed_f32x2( | ||
| xy_rounded, (fp32_round_int, fp32_round_int) | ||
| ) | ||
| - xy_frac = quack.activation.sub_packed_f32x2(xy_clamped, xy_rounded_back) | ||
| + xy_frac = cute.arch.sub_packed_f32x2(xy_clamped, xy_rounded_back) | ||
| xy_frac_ex2 = evaluate_polynomial_2(*xy_frac, POLY_EX2[poly_degree], loc=loc, ip=ip) | ||
| x_out = combine_int_frac_ex2(xy_rounded[0], xy_frac_ex2[0], loc=loc, ip=ip) | ||
| y_out = combine_int_frac_ex2(xy_rounded[1], xy_frac_ex2[1], loc=loc, ip=ip) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,33 @@ | ||
| Backport of https://github.com/state-spaces/mamba/commit/653923ce8fb0d47cdd9bcfd5904a0f1d58f91274 | ||
| Build the existing pinned Mamba sources with the standard required by PyTorch's ATen headers. | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Same here.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This patch fixes native-extension build failures because our pinned Mamba explicitly uses C++17 while the new PyTorch ATen headers require C++20; it changes only four compiler flags. No published Mamba release currently contains the fix, and state-spaces/mamba#1000 remains unmerged, so we retain the patch until a suitable upstream revision is adopted and validated.
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Can you please add this comment to the actual file? |
||
|
|
||
| diff --git a/setup.py b/setup.py | ||
| --- a/setup.py | ||
| +++ b/setup.py | ||
| @@ -210,10 +210,10 @@ def append_nvcc_threads(nvcc_extra_args): | ||
| if HIP_BUILD: | ||
|
|
||
| extra_compile_args = { | ||
| - "cxx": ["-O3", "-std=c++17"], | ||
| + "cxx": ["-O3", "-std=c++20"], | ||
| "nvcc": [ | ||
| "-O3", | ||
| - "-std=c++17", | ||
| + "-std=c++20", | ||
| f"--offload-arch={os.getenv('HIP_ARCHITECTURES', 'native')}", | ||
| "-U__CUDA_NO_HALF_OPERATORS__", | ||
| "-U__CUDA_NO_HALF_CONVERSIONS__", | ||
| @@ -223,11 +223,11 @@ def append_nvcc_threads(nvcc_extra_args): | ||
| } | ||
| else: | ||
| extra_compile_args = { | ||
| - "cxx": ["-O3", "-std=c++17"], | ||
| + "cxx": ["-O3", "-std=c++20"], | ||
| "nvcc": append_nvcc_threads( | ||
| [ | ||
| "-O3", | ||
| - "-std=c++17", | ||
| + "-std=c++20", | ||
| "-U__CUDA_NO_HALF_OPERATORS__", | ||
| "-U__CUDA_NO_HALF_CONVERSIONS__", | ||
| "-U__CUDA_NO_BFLOAT16_OPERATORS__", | ||
Uh oh!
There was an error while loading. Please reload this page.