-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patharith.patch
More file actions
97 lines (86 loc) · 4.13 KB
/
Copy patharith.patch
File metadata and controls
97 lines (86 loc) · 4.13 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
diff --git a/src/flag_gems/runtime/backend/_ascend/ops/pow.py b/src/flag_gems/runtime/backend/_ascend/ops/pow.py
index 5f5e3ef..ec386e0 100644
--- a/src/flag_gems/runtime/backend/_ascend/ops/pow.py
+++ b/src/flag_gems/runtime/backend/_ascend/ops/pow.py
@@ -1,4 +1,4 @@
-# Copyright 2026 FlagOS Contributors
+# Copyright 2026 FlagOS Contributors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@@ -24,15 +24,32 @@ _pow = tl_extra_shim.pow
logger = logging.getLogger(__name__)
+# wt-2026-09-15-fix (#5732): chained boolean operators (A or B or C) are not
+# supported by the triton-ascend 3.2 code generator that backends.yaml pins for
+# ascend-cann850 (flagtree 0.6.0+ascend3.2) — the kernel fails to compile with
+# UnsupportedLanguageConstruct. Parenthesized nesting (A or B) or C is accepted
+# by both 3.2 and 3.5 and is semantically identical, so we use it here.
+# wt <wangt635@ustc.edu.cn>
+#
+# wt-2026-09-15-fix (#5732 follow-up): the fp16/bf16-exponent branches passed
+# the exponent through unchanged while upcasting x to fp32, hitting CANN
+# libdevice pow's dispatch table which only has (fp32,fp32)/(fp16,fp16)/
+# (bf16,bf16) — KeyError(fp32, fp16), i.e. the same class as geluFIX #5867.
+# The old chained-or compile failure (triton 3.2) always masked this on pinned
+# environments. Upcast the exponent to fp32 as well: __hmf_powf is fp32 anyway,
+# numerics unchanged (verified bitwise on int/half/float exponent paths).
+
+
@pointwise_dynamic(promotion_methods=[(0, 1, "BOOL_TO_LONG")])
@triton.jit
def pow_func(x, exponent):
if (
- tl.constexpr(exponent.dtype.is_fp32())
- or tl.constexpr(exponent.dtype.is_fp16())
+ (tl.constexpr(exponent.dtype.is_fp32()) or tl.constexpr(exponent.dtype.is_fp16()))
or tl.constexpr(exponent.dtype.is_bf16())
):
- return _pow(x.to(tl.float32), exponent)
+ # exponent fp16/bf16 -> fp32 (fp32 命中表内 (fp32,fp32); 半精度原样传会
+ # 组成 (fp32, fp16) 查表 KeyError)
+ return _pow(x.to(tl.float32), exponent.to(tl.float32))
return _pow(x.to(tl.float32), exponent.to(tl.float32))
@@ -52,11 +69,10 @@ def pow_tensor_tensor_(A, exponent):
@triton.jit
def pow_func_tensor_scalar(x, exponent):
if (
- tl.constexpr(exponent.dtype.is_fp32())
- or tl.constexpr(exponent.dtype.is_fp16())
+ (tl.constexpr(exponent.dtype.is_fp32()) or tl.constexpr(exponent.dtype.is_fp16()))
or tl.constexpr(exponent.dtype.is_bf16())
):
- return _pow(x.to(tl.float32), exponent)
+ return _pow(x.to(tl.float32), exponent.to(tl.float32))
return _pow(x.to(tl.float32), exponent.to(tl.float32))
@@ -76,11 +92,10 @@ def pow_tensor_scalar_(A, exponent):
@triton.jit
def pow_func_scalar_tensor(x, exponent):
if (
- tl.constexpr(exponent.dtype.is_fp32())
- or tl.constexpr(exponent.dtype.is_fp16())
+ (tl.constexpr(exponent.dtype.is_fp32()) or tl.constexpr(exponent.dtype.is_fp16()))
or tl.constexpr(exponent.dtype.is_bf16())
):
- return _pow(x.to(tl.float32), exponent)
+ return _pow(x.to(tl.float32), exponent.to(tl.float32))
return _pow(x.to(tl.float32), exponent.to(tl.float32))
diff --git a/src/flag_gems/runtime/backend/_spacemit/ops/argmin.py b/src/flag_gems/runtime/backend/_spacemit/ops/argmin.py
index ab063e3..01f9d51 100644
--- a/src/flag_gems/runtime/backend/_spacemit/ops/argmin.py
+++ b/src/flag_gems/runtime/backend/_spacemit/ops/argmin.py
@@ -27,6 +27,9 @@ from flag_gems.utils.limits import get_dtype_max
logger = logging.getLogger(__name__)
+# wt-2026-09-15-fix (#5732): parenthesized chain for triton-ascend 3.2 compat (chained bool ops unsupported); semantically identical
+# wt <wangt635@ustc.edu.cn>
+
@libentry()
@triton.jit
@@ -86,7 +89,7 @@ def argmin_kernel(
if (dtype is tl.bfloat16 or dtype is tl.float16)
else (
tl.int32
- if (dtype is tl.int16 or dtype is tl.int8 or dtype is tl.uint8)
+ if ((dtype is tl.int16 or dtype is tl.int8) or dtype is tl.uint8)
else dtype
)
)