Repository navigation
Expand file tree
/
Copy pathpointwise_dynamic.patch
More file actions
91 lines (84 loc) · 4.42 KB
/
Copy pathpointwise_dynamic.patch
File metadata and controls
91 lines (84 loc) · 4.42 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
diff --git a/src/flag_gems/utils/pointwise_dynamic.py b/src/flag_gems/utils/pointwise_dynamic.py
index b499310..28efc08 100644
--- a/src/flag_gems/utils/pointwise_dynamic.py
+++ b/src/flag_gems/utils/pointwise_dynamic.py
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
+import dataclasses
import importlib
import os
from dataclasses import dataclass
@@ -1318,9 +1319,20 @@ class PointwiseDynamicFunction:
def _call_real_impl(self, *args, **kwargs):
"""Single entry point for real kernel invocation."""
- ndim, args, kwargs = self.prepare_args(*args, **kwargs)
- overload = self.instantiate(ndim)
- out = overload(*args, **kwargs)
+ ndim, args, kwargs, call_config = self.prepare_args(*args, **kwargs)
+ # wt-2026-09-11-fix (#5739): prepare_args may return a per-call config
+ # (block pointer disabled for broadcasted operands). Use it only for
+ # this call's instantiation so the shared vendor config is never
+ # polluted, and always restore afterwards (thread-safety best effort).
+ # wt <wangt635@ustc.edu.cn>
+ saved_config = self.config
+ if call_config is not None:
+ self.config = call_config
+ try:
+ overload = self.instantiate(ndim)
+ out = overload(*args, **kwargs)
+ finally:
+ self.config = saved_config
return self._unwrap(out)
# -------------------- complex helpers --------------------
@@ -1529,8 +1541,13 @@ class PointwiseDynamicFunction:
tensors = out_tensors + in_tensors
INT32_MAX = torch.iinfo(torch.int32).max
- if tensors[0].numel() > INT32_MAX:
- self.config.prefer_block_pointer = False
+ # wt-2026-09-11-fix (#5739): the old code mutated the *shared* vendor
+ # config in place (self.config.prefer_block_pointer = False), silently
+ # disabling block pointers for every other pointwise operator after one
+ # big call (singleton pollution, also called out by upstream PR #5668).
+ # wt <wangt635@ustc.edu.cn>
+ call_config = None
+ need_no_block_pointer = tensors[0].numel() > INT32_MAX
if self.use_fast_path(tensors): # dimension collapse & use physical ordering
allocated_outputs = [
torch.empty_like(tensors[0], dtype=dtype)
@@ -1563,6 +1580,23 @@ class PointwiseDynamicFunction:
task_shape = broadcast_shapes(shapes)
+ # wt-2026-09-11-fix (#5739, Ascend): broadcasted operands get a zero
+ # stride on expanded axes (broadcasted_stride), which block pointers
+ # cannot represent — the TritonToLinalg pipeline rejects them with
+ # "strides must not be zero" (ConvertTritonIRToLinalgIR on
+ # triton-ascend 3.5.1; ConvertLinalgRToBinary on the issue's 3.2).
+ # Reproduced: with prefer_block_pointer=True, non-broadcast add
+ # compiles fine while broadcast (128,1)+(1,128) always fails; the
+ # zero stride is the only variable. Disable block pointers for this
+ # call only (linear addressing handles stride 0 natively) and route
+ # through a copied config so the shared vendor singleton is untouched.
+ # wt <wangt635@ustc.edu.cn>
+ if not need_no_block_pointer:
+ for item in tensors:
+ if item.shape != task_shape:
+ need_no_block_pointer = True
+ break
+
if out_tensors:
for index, item in enumerate(out_tensors):
if list(item.shape) != list(task_shape):
@@ -1616,7 +1650,13 @@ class PointwiseDynamicFunction:
task_shape,
broadcasted_stride(item.shape, item.stride(), task_shape),
)
- return (ndim, args, kwargs)
+ # wt-2026-09-11-fix (#5739): materialize the per-call config once here;
+ # dataclasses.replace copies the shared vendor config so flipping
+ # prefer_block_pointer off never leaks to other operators.
+ # wt <wangt635@ustc.edu.cn>
+ if need_no_block_pointer and self.config.prefer_block_pointer:
+ call_config = dataclasses.replace(self.config, prefer_block_pointer=False)
+ return (ndim, args, kwargs, call_config)
def _unwrap(self, tensors):
# unwrap StridedBuffer to get Tensor