-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcopy.patch
More file actions
80 lines (66 loc) · 3.4 KB
/
Copy pathcopy.patch
File metadata and controls
80 lines (66 loc) · 3.4 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
diff --git a/src/flag_gems/runtime/backend/_kunlunxin/ops/copy.py b/src/flag_gems/runtime/backend/_kunlunxin/ops/copy.py
index c41c093c5cd13c127f2cd97102691ac5c79d292f..ad21ba2c217b001702ea36383397b22742b2ba9c 100644
--- a/src/flag_gems/runtime/backend/_kunlunxin/ops/copy.py
+++ b/src/flag_gems/runtime/backend/_kunlunxin/ops/copy.py
@@ -4,8 +4,8 @@ from typing import Optional
import torch
import triton
-from ..utils.codegen_config_utils import CodeGenConfig
-from ..utils.pointwise_dynamic import pointwise_dynamic
+from _kunlunxin.utils.codegen_config_utils import CodeGenConfig
+from _kunlunxin.utils.pointwise_dynamic import pointwise_dynamic
logger = logging.getLogger("flag_gems").getChild(__name__.lstrip("."))
@@ -43,6 +43,9 @@ def _copy_kernel(src):
return src
+# wt-2026-09-11-perf <wangt635@ustc.edu.cn>: non-contiguous sources are valid
+# pointwise inputs. Keep them on the tuned generated kernel instead of falling
+# back to the broken native path or a slower hand-written fixed-rank kernel.
def _can_use_triton(dst: torch.Tensor, src: torch.Tensor) -> bool:
if dst.layout != torch.strided or src.layout != torch.strided:
return False
@@ -54,11 +57,13 @@ def _can_use_triton(dst: torch.Tensor, src: torch.Tensor) -> bool:
# Preserve PyTorch's behaviour of warning when casting complex to real
# by forcing the redispatch path, which issues the warning internally.
return False
- if not src.is_contiguous():
+ if any(size > 1 and stride == 0 for size, stride in zip(dst.shape, dst.stride())):
+ # PyTorch rejects writes to an internally overlapping destination.
return False
return True
+# wt-2026-09-09-fix <wangt635@ustc.edu.cn>: reject overlapping writes before Triton.
def _expand_like(src: torch.Tensor, target_shape: torch.Size) -> torch.Tensor:
if src.shape == target_shape:
return src
@@ -77,7 +82,9 @@ def copy(
def copy_(dst: torch.Tensor, src: torch.Tensor, non_blocking: bool = False):
- if not isinstance(src, torch.Tensor):
+ if isinstance(src, (int, float, bool)):
+ src = torch.tensor(src, device=dst.device)
+ elif not isinstance(src, torch.Tensor):
raise TypeError("src must be a Tensor")
# this is the same as PyTorch's check
@@ -108,12 +115,6 @@ def copy_(dst: torch.Tensor, src: torch.Tensor, non_blocking: bool = False):
_FALLBACK_KEYSET, dst, src, non_blocking
)
- if dst.numel() == 0:
- # Respect PyTorch behaviour: empty tensors should still validate broadcast.
- return torch.ops.aten.copy_.default.redispatch(
- _FALLBACK_KEYSET, dst, src, non_blocking
- )
-
logger.debug("GEMS COPY_")
try:
@@ -126,8 +127,15 @@ def copy_(dst: torch.Tensor, src: torch.Tensor, non_blocking: bool = False):
f"The broadcast shape {broadcast_shape} does not match destination shape {tuple(dst.shape)}"
)
+ if dst.numel() == 0:
+ # Broadcast compatibility has already been checked above.
+ return dst
+
expanded_src = _expand_like(src, dst.shape)
+ # The generated pointwise kernel already carries true source/destination
+ # strides and uses the KunlunXIN-tuned grid/tile policy. Benchmarks on P800
+ # show it is substantially faster than launching the fixed-rank fallback.
overload = _copy_kernel.instantiate(expanded_src.ndim)
overload(expanded_src, out0=dst)
return dst