MusYW commited on
Commit
5b3773a
·
verified ·
1 Parent(s): 7f0c6e6

Training in progress, step 680

Browse files
Files changed (22) hide show
  1. .gitattributes +3 -0
  2. model.safetensors +1 -1
  3. torchinductor_ch-epfl-345354-j/aotautograd/akum3jin35di7tlhk4tvtzrbeforlbveooabshfrmtcm4cwst6h3/entry +0 -0
  4. torchinductor_ch-epfl-345354-j/aotautograd/azlrpmzq25mv3r2ycoghvuwp5ap446dq3pwd3e73xi4lravn37ws/entry +0 -0
  5. torchinductor_ch-epfl-345354-j/bw/cbw2ylnzyqqvpuutlhdcnot22eonl4ebucgx7hpqh66fjy7zgkhb.py +361 -0
  6. torchinductor_ch-epfl-345354-j/da/0df06b8658127e5ce915248b1384d7c58762d116d843ac58b2d4be91c249301b.best_config +1 -0
  7. torchinductor_ch-epfl-345354-j/da/cdapkyn6r3i3fmlikyctcbjr336rmrrljtxcaqcwrc7auok6odb4.py +46 -0
  8. torchinductor_ch-epfl-345354-j/fj/cfjpgg4ctgzyci7wzylrvpq3i7o2o5r6brzgt5u6bzfhltzrjid3.py +48 -0
  9. torchinductor_ch-epfl-345354-j/fj/fa90c6e1e9df4537c3b4e7a14a704c2df4c98a82068ce1f3abb21af8a35247a0.best_config +1 -0
  10. torchinductor_ch-epfl-345354-j/fxgraph/3o/f3odjl3g6dddsqtuut46oufxq75twuknijzu35ozpse6s2ndf7rh/ssax6qk4l2jsc4ovqf64hhenn3rpvkiqhluzs7qfxcc6dzbigdx +3 -0
  11. torchinductor_ch-epfl-345354-j/fxgraph/7y/f7yf7bpzpliuud5ylt3deq7gx7odarlyvhr2zdhcorkmqjxqlpv7/inrn63eoovhcqnqsled5nyfdmzwunnsaud5rcn37go5egvyuze3 +3 -0
  12. torchinductor_ch-epfl-345354-j/fxgraph/p3/fp32tea4g6rsybgfroshvlbxseks7iiqpb4cmbb2d3lgralck774/vbrtzv3sqiuroqkjniy4wufikxjgiptmxb4qwrykjuzdrp2mvvt +3 -0
  13. torchinductor_ch-epfl-345354-j/j6/42b81ae014f6f1c38797163131fc53a645adf2cc0424a9255831be278e34af78.best_config +1 -0
  14. torchinductor_ch-epfl-345354-j/j6/cj6yl63qj32gvorczxkszjnojwnybujlaozjrnzwzik4ukxq4rss.py +48 -0
  15. torchinductor_ch-epfl-345354-j/lm/65cc8fd7bae3abb18e0f1b23c8e1c60711fe732716b5724936da88bfeb6409b4.best_config +1 -0
  16. torchinductor_ch-epfl-345354-j/lm/clmeyn6qnatdy2hjtvp2smmgjoj5zqvy23kyx4cn7bbcafehg7wr.py +46 -0
  17. torchinductor_ch-epfl-345354-j/mc/cmcqiflrggybh55qgleeqzkbj3rlpmylxaqgl47omwv5z3y4ifv5.py +353 -0
  18. torchinductor_ch-epfl-345354-j/sw/cswtunku7iwygay5azmeyjzubvchodclgrov3ytwoqfke3cbjeek.py +357 -0
  19. torchinductor_ch-epfl-345354-j/sy/4d7e479de298eba85b0b0b838eb227f9e0129c56097eee3ae7375c0bd11c1c7d.best_config +1 -0
  20. torchinductor_ch-epfl-345354-j/sy/csyofljuv4nybcwnbyfhmdfufvm6xfyypzxkd745lbsexq35kdpk.py +46 -0
  21. torchinductor_ch-epfl-345354-j/ua/5c2b35108d0f6d838f30c1b016c317aa650f5f275110f17a9aff57ce54e364e6.best_config +1 -0
  22. torchinductor_ch-epfl-345354-j/ua/cuaufinvbemnicqkymrm4ujjq5b6ltwjbdc4jcgwj2zxj4usa5kx.py +46 -0
.gitattributes CHANGED
@@ -52,3 +52,6 @@ torchinductor_ch-epfl-345354-j/fxgraph/xb/fxbygtwj5f3jfikmhfoz7vd3mtceabjyw5ggsc
52
  torchinductor_ch-epfl-345354-j/triton/0/35ZDQUBPR56EFY64BDG6UB7OKJODSPJPYCTXXIKSOWZ2CA3EPWGA/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
53
  torchinductor_ch-epfl-345354-j/triton/0/SS6ERMBZYPQUFPYGWYCICDOLY35WYNXDIHS3C45P5EZJOFLNB5SQ/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
54
  torchinductor_ch-epfl-345354-j/triton/0/YGRXAKJ5T6UCGDY6RY3GMSSD4JVJFIRM6WYZVP7IT63SMM7NWHSA/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
 
 
 
 
52
  torchinductor_ch-epfl-345354-j/triton/0/35ZDQUBPR56EFY64BDG6UB7OKJODSPJPYCTXXIKSOWZ2CA3EPWGA/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
53
  torchinductor_ch-epfl-345354-j/triton/0/SS6ERMBZYPQUFPYGWYCICDOLY35WYNXDIHS3C45P5EZJOFLNB5SQ/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
54
  torchinductor_ch-epfl-345354-j/triton/0/YGRXAKJ5T6UCGDY6RY3GMSSD4JVJFIRM6WYZVP7IT63SMM7NWHSA/triton_poi_fused__to_copy_cos_mul_sin_1.cubin filter=lfs diff=lfs merge=lfs -text
55
+ torchinductor_ch-epfl-345354-j/fxgraph/3o/f3odjl3g6dddsqtuut46oufxq75twuknijzu35ozpse6s2ndf7rh/ssax6qk4l2jsc4ovqf64hhenn3rpvkiqhluzs7qfxcc6dzbigdx filter=lfs diff=lfs merge=lfs -text
56
+ torchinductor_ch-epfl-345354-j/fxgraph/7y/f7yf7bpzpliuud5ylt3deq7gx7odarlyvhr2zdhcorkmqjxqlpv7/inrn63eoovhcqnqsled5nyfdmzwunnsaud5rcn37go5egvyuze3 filter=lfs diff=lfs merge=lfs -text
57
+ torchinductor_ch-epfl-345354-j/fxgraph/p3/fp32tea4g6rsybgfroshvlbxseks7iiqpb4cmbb2d3lgralck774/vbrtzv3sqiuroqkjniy4wufikxjgiptmxb4qwrykjuzdrp2mvvt filter=lfs diff=lfs merge=lfs -text
model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:26045c44c8f7b94239544f69255de30ea5e05c7f8e23f0c01a67e755eaa0beba
3
  size 1192135096
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:17e7ee8610ecc91fe729f86ae30fbe910c6bf626a6ee7655d8d6496c22403797
3
  size 1192135096
torchinductor_ch-epfl-345354-j/aotautograd/akum3jin35di7tlhk4tvtzrbeforlbveooabshfrmtcm4cwst6h3/entry ADDED
Binary file (26.9 kB). View file
 
torchinductor_ch-epfl-345354-j/aotautograd/azlrpmzq25mv3r2ycoghvuwp5ap446dq3pwd3e73xi4lravn37ws/entry ADDED
Binary file (7.82 kB). View file
 
torchinductor_ch-epfl-345354-j/bw/cbw2ylnzyqqvpuutlhdcnot22eonl4ebucgx7hpqh66fjy7zgkhb.py ADDED
@@ -0,0 +1,361 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Compile-time auto-tuning block:
3
+
4
+ import torch
5
+ from torch._dynamo.testing import rand_strided
6
+ from torch._dynamo.utils import preserve_rng_state
7
+ from torch._inductor.select_algorithm import AlgorithmSelectorCache
8
+ from torch._inductor.async_compile import AsyncCompile
9
+
10
+ async_compile = AsyncCompile()
11
+ generate_example_value = AlgorithmSelectorCache.generate_example_value
12
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
13
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
14
+
15
+
16
+ triton_poi_fused_add_mul_neg_slice_backward_0 = async_compile.triton('triton_poi_fused_add_mul_neg_slice_backward_0', '''
17
+ import triton
18
+ import triton.language as tl
19
+
20
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
21
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
22
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
23
+ triton_helpers.set_driver_to_gpu()
24
+
25
+ @triton_heuristics.pointwise(
26
+ size_hints={'x': 1048576},
27
+ filename=__file__,
28
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
29
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
30
+ min_elem_per_thread=0
31
+ )
32
+ @triton.jit
33
+ def triton_poi_fused_add_mul_neg_slice_backward_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
34
+ xoffset = tl.program_id(0) * XBLOCK
35
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
36
+ xmask = xindex < xnumel
37
+ x0 = (xindex % ks0)
38
+ x3 = xindex
39
+ x4 = (xindex % ks1)
40
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
41
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
42
+ tmp0 = x0
43
+ tmp1 = ks0 // 2
44
+ tmp2 = tmp0 >= tmp1
45
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
46
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
47
+ tmp5 = tmp3 * tmp4
48
+ tmp6 = -tmp5
49
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
50
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
51
+ tmp9 = 0.0
52
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
53
+ tmp11 = tmp0 < tmp1
54
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
55
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
56
+ tmp14 = tmp12 * tmp13
57
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
58
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
59
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
60
+ tmp18 = tmp10 + tmp17
61
+ tmp21 = tmp19 * tmp20
62
+ tmp22 = tmp18 + tmp21
63
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
64
+ ''', device_str='cuda')
65
+
66
+
67
+ triton_poi_fused_add_mul_neg_slice_backward_1 = async_compile.triton('triton_poi_fused_add_mul_neg_slice_backward_1', '''
68
+ import triton
69
+ import triton.language as tl
70
+
71
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
72
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
73
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
74
+ triton_helpers.set_driver_to_gpu()
75
+
76
+ @triton_heuristics.pointwise(
77
+ size_hints={'x': 2097152},
78
+ filename=__file__,
79
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
80
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
81
+ min_elem_per_thread=0
82
+ )
83
+ @triton.jit
84
+ def triton_poi_fused_add_mul_neg_slice_backward_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
85
+ xoffset = tl.program_id(0) * XBLOCK
86
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
87
+ xmask = xindex < xnumel
88
+ x0 = (xindex % ks0)
89
+ x3 = xindex
90
+ x4 = (xindex % ks1)
91
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
92
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
93
+ tmp0 = x0
94
+ tmp1 = ks0 // 2
95
+ tmp2 = tmp0 >= tmp1
96
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
97
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
98
+ tmp5 = tmp3 * tmp4
99
+ tmp6 = -tmp5
100
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
101
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
102
+ tmp9 = 0.0
103
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
104
+ tmp11 = tmp0 < tmp1
105
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
106
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
107
+ tmp14 = tmp12 * tmp13
108
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
109
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
110
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
111
+ tmp18 = tmp10 + tmp17
112
+ tmp21 = tmp19 * tmp20
113
+ tmp22 = tmp18 + tmp21
114
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
115
+ ''', device_str='cuda')
116
+
117
+ async_compile.wait(globals())
118
+ del async_compile
119
+
120
+ import triton
121
+ import triton.language as tl
122
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
123
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
124
+ with torch.cuda._DeviceGuard(0):
125
+ torch.cuda.set_device(0)
126
+ stream0 = get_raw_stream(0)
127
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
128
+ stream0 = get_raw_stream(0)
129
+ tangents_2 = generate_example_value((4, 8, 177, 128), (181248, 22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
130
+ primals_5 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
131
+ primals_3 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
132
+ buf0 = generate_example_value((4, 8, 177, 128), (181248, 22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
133
+ triton_poi_fused_add_mul_neg_slice_backward_0.run(tangents_2, primals_5, primals_3, buf0, 128, 22656, 724992, stream=stream0)
134
+ del tangents_2, primals_5, primals_3, buf0
135
+
136
+ stream0 = get_raw_stream(0)
137
+ tangents_1 = generate_example_value((4, 16, 177, 128), (362496, 22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
138
+ primals_5 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
139
+ primals_3 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
140
+ buf1 = generate_example_value((4, 16, 177, 128), (362496, 22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
141
+ triton_poi_fused_add_mul_neg_slice_backward_1.run(tangents_1, primals_5, primals_3, buf1, 128, 22656, 1449984, stream=stream0)
142
+ del tangents_1, primals_5, primals_3, buf1
143
+
144
+ """
145
+ # AOT ID: ['13_backward']
146
+ from ctypes import c_void_p, c_long, c_int
147
+ import torch
148
+ import math
149
+ import random
150
+ import os
151
+ import tempfile
152
+ from math import inf, nan
153
+ from cmath import nanj
154
+ from torch._inductor.hooks import run_intermediate_hooks
155
+ from torch._inductor.utils import maybe_profile
156
+ from torch._inductor.codegen.memory_planning import _align as align
157
+ from torch import device, empty_strided
158
+ from torch._inductor.async_compile import AsyncCompile
159
+ from torch._inductor.select_algorithm import extern_kernels
160
+ from torch._inductor.codegen.multi_kernel import MultiKernelCall
161
+ import triton
162
+ import triton.language as tl
163
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
164
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
165
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
166
+
167
+ aten = torch.ops.aten
168
+ inductor_ops = torch.ops.inductor
169
+ _quantized = torch.ops._quantized
170
+ assert_size_stride = torch._C._dynamo.guards.assert_size_stride
171
+ empty_strided_cpu = torch._C._dynamo.guards._empty_strided_cpu
172
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
173
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
174
+ reinterpret_tensor = torch._C._dynamo.guards._reinterpret_tensor
175
+ alloc_from_pool = torch.ops.inductor._alloc_from_pool
176
+ async_compile = AsyncCompile()
177
+ empty_strided_p2p = torch._C._distributed_c10d._SymmetricMemory.empty_strided_p2p
178
+
179
+
180
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/j6/cj6yl63qj32gvorczxkszjnojwnybujlaozjrnzwzik4ukxq4rss.py
181
+ # Topologically Sorted Source Nodes: [], Original ATen: [aten.neg, aten.slice_backward, aten.add, aten.mul]
182
+ # Source node to ATen node mapping:
183
+ # Graph fragment:
184
+ # %neg_2 : [num_users=1] = call_function[target=torch.ops.aten.neg.default](args = (%slice_5,), kwargs = {})
185
+ # %full_default : [num_users=2] = call_function[target=torch.ops.aten.full.default](args = ([%primals_9, %primals_10, %primals_1, %primals_2], 0), kwargs = {dtype: torch.bfloat16, layout: torch.strided, device: cuda:0, pin_memory: False})
186
+ # %slice_scatter_default : [num_users=1] = call_function[target=torch.ops.aten.slice_scatter.default](args = (%full_default, %neg_2, 3, %floordiv, 9223372036854775807), kwargs = {})
187
+ # %slice_scatter_default_1 : [num_users=1] = call_function[target=torch.ops.aten.slice_scatter.default](args = (%full_default, %slice_6, 3, 0, %floordiv), kwargs = {})
188
+ # %add_82 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%slice_scatter_default, %slice_scatter_default_1), kwargs = {})
189
+ # %mul_69 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%tangents_2, %unsqueeze), kwargs = {})
190
+ # %add_83 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%add_82, %mul_69), kwargs = {})
191
+ triton_poi_fused_add_mul_neg_slice_backward_0 = async_compile.triton('triton_poi_fused_add_mul_neg_slice_backward_0', '''
192
+ import triton
193
+ import triton.language as tl
194
+
195
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
196
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
197
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
198
+ triton_helpers.set_driver_to_gpu()
199
+
200
+ @triton_heuristics.pointwise(
201
+ size_hints={'x': 1048576},
202
+ filename=__file__,
203
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
204
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
205
+ min_elem_per_thread=0
206
+ )
207
+ @triton.jit
208
+ def triton_poi_fused_add_mul_neg_slice_backward_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
209
+ xoffset = tl.program_id(0) * XBLOCK
210
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
211
+ xmask = xindex < xnumel
212
+ x0 = (xindex % ks0)
213
+ x3 = xindex
214
+ x4 = (xindex % ks1)
215
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
216
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
217
+ tmp0 = x0
218
+ tmp1 = ks0 // 2
219
+ tmp2 = tmp0 >= tmp1
220
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
221
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
222
+ tmp5 = tmp3 * tmp4
223
+ tmp6 = -tmp5
224
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
225
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
226
+ tmp9 = 0.0
227
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
228
+ tmp11 = tmp0 < tmp1
229
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
230
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
231
+ tmp14 = tmp12 * tmp13
232
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
233
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
234
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
235
+ tmp18 = tmp10 + tmp17
236
+ tmp21 = tmp19 * tmp20
237
+ tmp22 = tmp18 + tmp21
238
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
239
+ ''', device_str='cuda')
240
+
241
+
242
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/fj/cfjpgg4ctgzyci7wzylrvpq3i7o2o5r6brzgt5u6bzfhltzrjid3.py
243
+ # Topologically Sorted Source Nodes: [], Original ATen: [aten.neg, aten.slice_backward, aten.add, aten.mul]
244
+ # Source node to ATen node mapping:
245
+ # Graph fragment:
246
+ # %neg_3 : [num_users=1] = call_function[target=torch.ops.aten.neg.default](args = (%slice_7,), kwargs = {})
247
+ # %full_default_2 : [num_users=2] = call_function[target=torch.ops.aten.full.default](args = ([%primals_6, %primals_7, %primals_1, %primals_2], 0), kwargs = {dtype: torch.bfloat16, layout: torch.strided, device: cuda:0, pin_memory: False})
248
+ # %slice_scatter_default_2 : [num_users=1] = call_function[target=torch.ops.aten.slice_scatter.default](args = (%full_default_2, %neg_3, 3, %floordiv, 9223372036854775807), kwargs = {})
249
+ # %slice_scatter_default_3 : [num_users=1] = call_function[target=torch.ops.aten.slice_scatter.default](args = (%full_default_2, %slice_8, 3, 0, %floordiv), kwargs = {})
250
+ # %add_88 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%slice_scatter_default_2, %slice_scatter_default_3), kwargs = {})
251
+ # %mul_71 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%tangents_1, %unsqueeze), kwargs = {})
252
+ # %add_89 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%add_88, %mul_71), kwargs = {})
253
+ triton_poi_fused_add_mul_neg_slice_backward_1 = async_compile.triton('triton_poi_fused_add_mul_neg_slice_backward_1', '''
254
+ import triton
255
+ import triton.language as tl
256
+
257
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
258
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
259
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
260
+ triton_helpers.set_driver_to_gpu()
261
+
262
+ @triton_heuristics.pointwise(
263
+ size_hints={'x': 2097152},
264
+ filename=__file__,
265
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
266
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
267
+ min_elem_per_thread=0
268
+ )
269
+ @triton.jit
270
+ def triton_poi_fused_add_mul_neg_slice_backward_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
271
+ xoffset = tl.program_id(0) * XBLOCK
272
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
273
+ xmask = xindex < xnumel
274
+ x0 = (xindex % ks0)
275
+ x3 = xindex
276
+ x4 = (xindex % ks1)
277
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
278
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
279
+ tmp0 = x0
280
+ tmp1 = ks0 // 2
281
+ tmp2 = tmp0 >= tmp1
282
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
283
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
284
+ tmp5 = tmp3 * tmp4
285
+ tmp6 = -tmp5
286
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
287
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
288
+ tmp9 = 0.0
289
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
290
+ tmp11 = tmp0 < tmp1
291
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
292
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
293
+ tmp14 = tmp12 * tmp13
294
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
295
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
296
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
297
+ tmp18 = tmp10 + tmp17
298
+ tmp21 = tmp19 * tmp20
299
+ tmp22 = tmp18 + tmp21
300
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
301
+ ''', device_str='cuda')
302
+
303
+
304
+ async_compile.wait(globals())
305
+ del async_compile
306
+
307
+ def call(args):
308
+ primals_1, primals_2, primals_6, primals_7, primals_9, primals_10, floordiv, add_78, primals_3, primals_5, tangents_1, tangents_2 = args
309
+ args.clear()
310
+ s0 = primals_1
311
+ s1 = primals_2
312
+ s7 = primals_6
313
+ s8 = primals_7
314
+ s13 = primals_9
315
+ s14 = primals_10
316
+ assert_size_stride(primals_3, (1, s0, s1), (s0*s1, s1, 1))
317
+ assert_size_stride(primals_5, (1, s0, s1), (s0*s1, s1, 1))
318
+ assert_size_stride(tangents_1, (s7, s8, s0, s1), (s0*s1*s8, s0*s1, s1, 1))
319
+ assert_size_stride(tangents_2, (s13, s14, s0, s1), (s0*s1*s14, s0*s1, s1, 1))
320
+ with torch.cuda._DeviceGuard(0):
321
+ torch.cuda.set_device(0)
322
+ ps0 = s0*s1
323
+ buf0 = empty_strided_cuda((s13, s14, s0, s1), (s0*s1*s14, s0*s1, s1, 1), torch.bfloat16)
324
+ # Topologically Sorted Source Nodes: [], Original ATen: [aten.neg, aten.slice_backward, aten.add, aten.mul]
325
+ triton_poi_fused_add_mul_neg_slice_backward_0_xnumel = s0*s1*s13*s14
326
+ stream0 = get_raw_stream(0)
327
+ triton_poi_fused_add_mul_neg_slice_backward_0.run(tangents_2, primals_5, primals_3, buf0, s1, ps0, triton_poi_fused_add_mul_neg_slice_backward_0_xnumel, stream=stream0)
328
+ del tangents_2
329
+ buf1 = empty_strided_cuda((s7, s8, s0, s1), (s0*s1*s8, s0*s1, s1, 1), torch.bfloat16)
330
+ # Topologically Sorted Source Nodes: [], Original ATen: [aten.neg, aten.slice_backward, aten.add, aten.mul]
331
+ triton_poi_fused_add_mul_neg_slice_backward_1_xnumel = s0*s1*s7*s8
332
+ stream0 = get_raw_stream(0)
333
+ triton_poi_fused_add_mul_neg_slice_backward_1.run(tangents_1, primals_5, primals_3, buf1, s1, ps0, triton_poi_fused_add_mul_neg_slice_backward_1_xnumel, stream=stream0)
334
+ del primals_3
335
+ del primals_5
336
+ del tangents_1
337
+ return (None, None, None, None, None, None, None, buf1, None, None, buf0, )
338
+
339
+
340
+ def benchmark_compiled_module(times=10, repeat=10):
341
+ from torch._dynamo.testing import rand_strided
342
+ from torch._inductor.utils import print_performance
343
+ primals_1 = 177
344
+ primals_2 = 128
345
+ primals_6 = 4
346
+ primals_7 = 16
347
+ primals_9 = 4
348
+ primals_10 = 8
349
+ floordiv = 64
350
+ add_78 = 64
351
+ primals_3 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
352
+ primals_5 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
353
+ tangents_1 = rand_strided((4, 16, 177, 128), (362496, 22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
354
+ tangents_2 = rand_strided((4, 8, 177, 128), (181248, 22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
355
+ fn = lambda: call([primals_1, primals_2, primals_6, primals_7, primals_9, primals_10, floordiv, add_78, primals_3, primals_5, tangents_1, tangents_2])
356
+ return print_performance(fn, times=times, repeat=repeat)
357
+
358
+
359
+ if __name__ == "__main__":
360
+ from torch._inductor.wrapper_benchmark import compiled_module_main
361
+ compiled_module_main('None', benchmark_compiled_module)
torchinductor_ch-epfl-345354-j/da/0df06b8658127e5ce915248b1384d7c58762d116d843ac58b2d4be91c249301b.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 1024, "num_warps": 4, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 25}
torchinductor_ch-epfl-345354-j/da/cdapkyn6r3i3fmlikyctcbjr336rmrrljtxcaqcwrc7auok6odb4.py ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 1048576},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x4 = xindex
23
+ x0 = (xindex % ks0)
24
+ x2 = ((xindex // ks1) % ks2)
25
+ x5 = xindex // ks0
26
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
28
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
29
+ tmp2 = tmp0 * tmp1
30
+ tmp3 = x0
31
+ tmp4 = tl.full([1], 0, tl.int64)
32
+ tmp5 = tmp3 >= tmp4
33
+ tmp6 = ks0 + (-1)*(ks0 // 2)
34
+ tmp7 = tmp3 < tmp6
35
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
36
+ tmp9 = -tmp8
37
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
38
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
39
+ tmp12 = tmp3 >= tmp6
40
+ tmp13 = ks0
41
+ tmp14 = tmp3 < tmp13
42
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
43
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
44
+ tmp18 = tmp16 * tmp17
45
+ tmp19 = tmp2 + tmp18
46
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
torchinductor_ch-epfl-345354-j/fj/cfjpgg4ctgzyci7wzylrvpq3i7o2o5r6brzgt5u6bzfhltzrjid3.py ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 2097152},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_mul_neg_slice_backward_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x0 = (xindex % ks0)
23
+ x3 = xindex
24
+ x4 = (xindex % ks1)
25
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
26
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp0 = x0
28
+ tmp1 = ks0 // 2
29
+ tmp2 = tmp0 >= tmp1
30
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
31
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
32
+ tmp5 = tmp3 * tmp4
33
+ tmp6 = -tmp5
34
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
35
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
36
+ tmp9 = 0.0
37
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
38
+ tmp11 = tmp0 < tmp1
39
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
40
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
41
+ tmp14 = tmp12 * tmp13
42
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
43
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
44
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
45
+ tmp18 = tmp10 + tmp17
46
+ tmp21 = tmp19 * tmp20
47
+ tmp22 = tmp18 + tmp21
48
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
torchinductor_ch-epfl-345354-j/fj/fa90c6e1e9df4537c3b4e7a14a704c2df4c98a82068ce1f3abb21af8a35247a0.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 512, "num_warps": 8, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 51}
torchinductor_ch-epfl-345354-j/fxgraph/3o/f3odjl3g6dddsqtuut46oufxq75twuknijzu35ozpse6s2ndf7rh/ssax6qk4l2jsc4ovqf64hhenn3rpvkiqhluzs7qfxcc6dzbigdx ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:94817f415c5e932171d5817dc39c8d6ee2faa9103ae9997e05b885e1e42830ee
3
+ size 599455
torchinductor_ch-epfl-345354-j/fxgraph/7y/f7yf7bpzpliuud5ylt3deq7gx7odarlyvhr2zdhcorkmqjxqlpv7/inrn63eoovhcqnqsled5nyfdmzwunnsaud5rcn37go5egvyuze3 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4362df6c8e754e2836125907d6e0a3666d46b640a0fb11377f33ba435714c937
3
+ size 600150
torchinductor_ch-epfl-345354-j/fxgraph/p3/fp32tea4g6rsybgfroshvlbxseks7iiqpb4cmbb2d3lgralck774/vbrtzv3sqiuroqkjniy4wufikxjgiptmxb4qwrykjuzdrp2mvvt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a8633cd77282291741496a31cb50a855d2643e6cb8790b470a4d3238bf4cad66
3
+ size 517209
torchinductor_ch-epfl-345354-j/j6/42b81ae014f6f1c38797163131fc53a645adf2cc0424a9255831be278e34af78.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 512, "num_warps": 8, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 51}
torchinductor_ch-epfl-345354-j/j6/cj6yl63qj32gvorczxkszjnojwnybujlaozjrnzwzik4ukxq4rss.py ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 1048576},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_mul_neg_slice_backward_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 6, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_mul_neg_slice_backward_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x0 = (xindex % ks0)
23
+ x3 = xindex
24
+ x4 = (xindex % ks1)
25
+ tmp19 = tl.load(in_ptr0 + (x3), xmask, eviction_policy='evict_last').to(tl.float32)
26
+ tmp20 = tl.load(in_ptr2 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp0 = x0
28
+ tmp1 = ks0 // 2
29
+ tmp2 = tmp0 >= tmp1
30
+ tmp3 = tl.load(in_ptr0 + (x3 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
31
+ tmp4 = tl.load(in_ptr1 + (x4 + (-1)*(ks0 // 2)), xmask & tmp2, eviction_policy='evict_last', other=0.0).to(tl.float32)
32
+ tmp5 = tmp3 * tmp4
33
+ tmp6 = -tmp5
34
+ tmp7 = tl.full(tmp6.shape, 0.0, tmp6.dtype)
35
+ tmp8 = tl.where(tmp2, tmp6, tmp7)
36
+ tmp9 = 0.0
37
+ tmp10 = tl.where(tmp2, tmp8, tmp9)
38
+ tmp11 = tmp0 < tmp1
39
+ tmp12 = tl.load(in_ptr0 + (ks0 + x3 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
40
+ tmp13 = tl.load(in_ptr1 + (ks0 + x4 + (-1)*(ks0 // 2)), xmask & tmp11, eviction_policy='evict_last', other=0.0).to(tl.float32)
41
+ tmp14 = tmp12 * tmp13
42
+ tmp15 = tl.full(tmp14.shape, 0.0, tmp14.dtype)
43
+ tmp16 = tl.where(tmp11, tmp14, tmp15)
44
+ tmp17 = tl.where(tmp11, tmp16, tmp9)
45
+ tmp18 = tmp10 + tmp17
46
+ tmp21 = tmp19 * tmp20
47
+ tmp22 = tmp18 + tmp21
48
+ tl.store(out_ptr0 + (x3), tmp22, xmask)
torchinductor_ch-epfl-345354-j/lm/65cc8fd7bae3abb18e0f1b23c8e1c60711fe732716b5724936da88bfeb6409b4.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 1024, "num_warps": 4, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 38}
torchinductor_ch-epfl-345354-j/lm/clmeyn6qnatdy2hjtvp2smmgjoj5zqvy23kyx4cn7bbcafehg7wr.py ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 1048576},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x4 = xindex
23
+ x0 = (xindex % ks0)
24
+ x2 = ((xindex // ks1) % ks2)
25
+ x5 = xindex // ks0
26
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
28
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
29
+ tmp2 = tmp0 * tmp1
30
+ tmp3 = x0
31
+ tmp4 = tl.full([1], 0, tl.int64)
32
+ tmp5 = tmp3 >= tmp4
33
+ tmp6 = ks0 + (-1)*(ks0 // 2)
34
+ tmp7 = tmp3 < tmp6
35
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
36
+ tmp9 = -tmp8
37
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
38
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
39
+ tmp12 = tmp3 >= tmp6
40
+ tmp13 = ks0
41
+ tmp14 = tmp3 < tmp13
42
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
43
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
44
+ tmp18 = tmp16 * tmp17
45
+ tmp19 = tmp2 + tmp18
46
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
torchinductor_ch-epfl-345354-j/mc/cmcqiflrggybh55qgleeqzkbj3rlpmylxaqgl47omwv5z3y4ifv5.py ADDED
@@ -0,0 +1,353 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Compile-time auto-tuning block:
3
+
4
+ import torch
5
+ from torch._dynamo.testing import rand_strided
6
+ from torch._dynamo.utils import preserve_rng_state
7
+ from torch._inductor.select_algorithm import AlgorithmSelectorCache
8
+ from torch._inductor.async_compile import AsyncCompile
9
+
10
+ async_compile = AsyncCompile()
11
+ generate_example_value = AlgorithmSelectorCache.generate_example_value
12
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
13
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
14
+
15
+
16
+ triton_poi_fused_add_cat_mul_0 = async_compile.triton('triton_poi_fused_add_cat_mul_0', '''
17
+ import triton
18
+ import triton.language as tl
19
+
20
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
21
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
22
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
23
+ triton_helpers.set_driver_to_gpu()
24
+
25
+ @triton_heuristics.pointwise(
26
+ size_hints={'x': 2097152},
27
+ filename=__file__,
28
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
29
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
30
+ min_elem_per_thread=0
31
+ )
32
+ @triton.jit
33
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
34
+ xoffset = tl.program_id(0) * XBLOCK
35
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
36
+ xmask = xindex < xnumel
37
+ x4 = xindex
38
+ x0 = (xindex % ks0)
39
+ x2 = ((xindex // ks1) % ks2)
40
+ x5 = xindex // ks0
41
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
42
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
43
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
44
+ tmp2 = tmp0 * tmp1
45
+ tmp3 = x0
46
+ tmp4 = tl.full([1], 0, tl.int64)
47
+ tmp5 = tmp3 >= tmp4
48
+ tmp6 = ks0 + (-1)*(ks0 // 2)
49
+ tmp7 = tmp3 < tmp6
50
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
51
+ tmp9 = -tmp8
52
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
53
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
54
+ tmp12 = tmp3 >= tmp6
55
+ tmp13 = ks0
56
+ tmp14 = tmp3 < tmp13
57
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
58
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
59
+ tmp18 = tmp16 * tmp17
60
+ tmp19 = tmp2 + tmp18
61
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
62
+ ''', device_str='cuda')
63
+
64
+
65
+ triton_poi_fused_add_cat_mul_1 = async_compile.triton('triton_poi_fused_add_cat_mul_1', '''
66
+ import triton
67
+ import triton.language as tl
68
+
69
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
70
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
71
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
72
+ triton_helpers.set_driver_to_gpu()
73
+
74
+ @triton_heuristics.pointwise(
75
+ size_hints={'x': 1048576},
76
+ filename=__file__,
77
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
78
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
79
+ min_elem_per_thread=0
80
+ )
81
+ @triton.jit
82
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
83
+ xoffset = tl.program_id(0) * XBLOCK
84
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
85
+ xmask = xindex < xnumel
86
+ x4 = xindex
87
+ x0 = (xindex % ks0)
88
+ x2 = ((xindex // ks1) % ks2)
89
+ x5 = xindex // ks0
90
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
91
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
92
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
93
+ tmp2 = tmp0 * tmp1
94
+ tmp3 = x0
95
+ tmp4 = tl.full([1], 0, tl.int64)
96
+ tmp5 = tmp3 >= tmp4
97
+ tmp6 = ks0 + (-1)*(ks0 // 2)
98
+ tmp7 = tmp3 < tmp6
99
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
100
+ tmp9 = -tmp8
101
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
102
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
103
+ tmp12 = tmp3 >= tmp6
104
+ tmp13 = ks0
105
+ tmp14 = tmp3 < tmp13
106
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
107
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
108
+ tmp18 = tmp16 * tmp17
109
+ tmp19 = tmp2 + tmp18
110
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
111
+ ''', device_str='cuda')
112
+
113
+ async_compile.wait(globals())
114
+ del async_compile
115
+
116
+ import triton
117
+ import triton.language as tl
118
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
119
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
120
+ with torch.cuda._DeviceGuard(0):
121
+ torch.cuda.set_device(0)
122
+ stream0 = get_raw_stream(0)
123
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
124
+ stream0 = get_raw_stream(0)
125
+ primals_8 = generate_example_value((4, 16, 177, 128), (362496, 128, 2048, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
126
+ primals_3 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
127
+ primals_5 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
128
+ buf0 = generate_example_value((4, 16, 177, 128), (362496, 128, 2048, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
129
+ triton_poi_fused_add_cat_mul_0.run(primals_8, primals_3, primals_5, buf0, 128, 2048, 177, 1449984, stream=stream0)
130
+ del primals_8, primals_3, primals_5, buf0
131
+
132
+ stream0 = get_raw_stream(0)
133
+ primals_11 = generate_example_value((4, 8, 177, 128), (181248, 128, 1024, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
134
+ primals_3 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
135
+ primals_5 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
136
+ buf1 = generate_example_value((4, 8, 177, 128), (181248, 128, 1024, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
137
+ triton_poi_fused_add_cat_mul_1.run(primals_11, primals_3, primals_5, buf1, 128, 1024, 177, 724992, stream=stream0)
138
+ del primals_11, primals_3, primals_5, buf1
139
+
140
+ """
141
+ # AOT ID: ['13_forward']
142
+ from ctypes import c_void_p, c_long, c_int
143
+ import torch
144
+ import math
145
+ import random
146
+ import os
147
+ import tempfile
148
+ from math import inf, nan
149
+ from cmath import nanj
150
+ from torch._inductor.hooks import run_intermediate_hooks
151
+ from torch._inductor.utils import maybe_profile
152
+ from torch._inductor.codegen.memory_planning import _align as align
153
+ from torch import device, empty_strided
154
+ from torch._inductor.async_compile import AsyncCompile
155
+ from torch._inductor.select_algorithm import extern_kernels
156
+ from torch._inductor.codegen.multi_kernel import MultiKernelCall
157
+ import triton
158
+ import triton.language as tl
159
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
160
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
161
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
162
+
163
+ aten = torch.ops.aten
164
+ inductor_ops = torch.ops.inductor
165
+ _quantized = torch.ops._quantized
166
+ assert_size_stride = torch._C._dynamo.guards.assert_size_stride
167
+ empty_strided_cpu = torch._C._dynamo.guards._empty_strided_cpu
168
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
169
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
170
+ reinterpret_tensor = torch._C._dynamo.guards._reinterpret_tensor
171
+ alloc_from_pool = torch.ops.inductor._alloc_from_pool
172
+ async_compile = AsyncCompile()
173
+ empty_strided_p2p = torch._C._distributed_c10d._SymmetricMemory.empty_strided_p2p
174
+
175
+
176
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/ua/cuaufinvbemnicqkymrm4ujjq5b6ltwjbdc4jcgwj2zxj4usa5kx.py
177
+ # Topologically Sorted Source Nodes: [mul, cat, mul_1, q_embed], Original ATen: [aten.mul, aten.cat, aten.add]
178
+ # Source node to ATen node mapping:
179
+ # cat => cat
180
+ # mul => mul_8
181
+ # mul_1 => mul_29
182
+ # q_embed => add_36
183
+ # Graph fragment:
184
+ # %mul_8 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%primals_8, %unsqueeze), kwargs = {})
185
+ # %cat : [num_users=1] = call_function[target=torch.ops.aten.cat.default](args = ([%neg, %slice_1], -1), kwargs = {})
186
+ # %mul_29 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%cat, %unsqueeze_1), kwargs = {})
187
+ # %add_36 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%mul_8, %mul_29), kwargs = {})
188
+ triton_poi_fused_add_cat_mul_0 = async_compile.triton('triton_poi_fused_add_cat_mul_0', '''
189
+ import triton
190
+ import triton.language as tl
191
+
192
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
193
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
194
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
195
+ triton_helpers.set_driver_to_gpu()
196
+
197
+ @triton_heuristics.pointwise(
198
+ size_hints={'x': 2097152},
199
+ filename=__file__,
200
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
201
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
202
+ min_elem_per_thread=0
203
+ )
204
+ @triton.jit
205
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
206
+ xoffset = tl.program_id(0) * XBLOCK
207
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
208
+ xmask = xindex < xnumel
209
+ x4 = xindex
210
+ x0 = (xindex % ks0)
211
+ x2 = ((xindex // ks1) % ks2)
212
+ x5 = xindex // ks0
213
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
214
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
215
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
216
+ tmp2 = tmp0 * tmp1
217
+ tmp3 = x0
218
+ tmp4 = tl.full([1], 0, tl.int64)
219
+ tmp5 = tmp3 >= tmp4
220
+ tmp6 = ks0 + (-1)*(ks0 // 2)
221
+ tmp7 = tmp3 < tmp6
222
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
223
+ tmp9 = -tmp8
224
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
225
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
226
+ tmp12 = tmp3 >= tmp6
227
+ tmp13 = ks0
228
+ tmp14 = tmp3 < tmp13
229
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
230
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
231
+ tmp18 = tmp16 * tmp17
232
+ tmp19 = tmp2 + tmp18
233
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
234
+ ''', device_str='cuda')
235
+
236
+
237
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/lm/clmeyn6qnatdy2hjtvp2smmgjoj5zqvy23kyx4cn7bbcafehg7wr.py
238
+ # Topologically Sorted Source Nodes: [mul_2, cat_1, mul_3, k_embed], Original ATen: [aten.mul, aten.cat, aten.add]
239
+ # Source node to ATen node mapping:
240
+ # cat_1 => cat_1
241
+ # k_embed => add_72
242
+ # mul_2 => mul_38
243
+ # mul_3 => mul_59
244
+ # Graph fragment:
245
+ # %mul_38 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%primals_11, %unsqueeze), kwargs = {})
246
+ # %cat_1 : [num_users=1] = call_function[target=torch.ops.aten.cat.default](args = ([%neg_1, %slice_3], -1), kwargs = {})
247
+ # %mul_59 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%cat_1, %unsqueeze_1), kwargs = {})
248
+ # %add_72 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%mul_38, %mul_59), kwargs = {})
249
+ triton_poi_fused_add_cat_mul_1 = async_compile.triton('triton_poi_fused_add_cat_mul_1', '''
250
+ import triton
251
+ import triton.language as tl
252
+
253
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
254
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
255
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
256
+ triton_helpers.set_driver_to_gpu()
257
+
258
+ @triton_heuristics.pointwise(
259
+ size_hints={'x': 1048576},
260
+ filename=__file__,
261
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
262
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
263
+ min_elem_per_thread=0
264
+ )
265
+ @triton.jit
266
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
267
+ xoffset = tl.program_id(0) * XBLOCK
268
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
269
+ xmask = xindex < xnumel
270
+ x4 = xindex
271
+ x0 = (xindex % ks0)
272
+ x2 = ((xindex // ks1) % ks2)
273
+ x5 = xindex // ks0
274
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
275
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
276
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
277
+ tmp2 = tmp0 * tmp1
278
+ tmp3 = x0
279
+ tmp4 = tl.full([1], 0, tl.int64)
280
+ tmp5 = tmp3 >= tmp4
281
+ tmp6 = ks0 + (-1)*(ks0 // 2)
282
+ tmp7 = tmp3 < tmp6
283
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
284
+ tmp9 = -tmp8
285
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
286
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
287
+ tmp12 = tmp3 >= tmp6
288
+ tmp13 = ks0
289
+ tmp14 = tmp3 < tmp13
290
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
291
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
292
+ tmp18 = tmp16 * tmp17
293
+ tmp19 = tmp2 + tmp18
294
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
295
+ ''', device_str='cuda')
296
+
297
+
298
+ async_compile.wait(globals())
299
+ del async_compile
300
+
301
+ def call(args):
302
+ primals_1, primals_2, primals_3, primals_4, primals_5, primals_6, primals_7, primals_8, primals_9, primals_10, primals_11 = args
303
+ args.clear()
304
+ s0 = primals_1
305
+ s1 = primals_2
306
+ s7 = primals_6
307
+ s8 = primals_7
308
+ s13 = primals_9
309
+ s14 = primals_10
310
+ assert_size_stride(primals_3, (1, s0, s1), (s0*s1, s1, 1))
311
+ assert_size_stride(primals_5, (1, s0, s1), (s0*s1, s1, 1))
312
+ assert_size_stride(primals_8, (s7, s8, s0, s1), (s0*s1*s8, s1, s1*s8, 1))
313
+ assert_size_stride(primals_11, (s13, s14, s0, s1), (s0*s1*s14, s1, s1*s14, 1))
314
+ with torch.cuda._DeviceGuard(0):
315
+ torch.cuda.set_device(0)
316
+ ps0 = s1*s8
317
+ buf0 = empty_strided_cuda((s7, s8, s0, s1), (s0*s1*s8, s1, s1*s8, 1), torch.bfloat16)
318
+ # Topologically Sorted Source Nodes: [mul, cat, mul_1, q_embed], Original ATen: [aten.mul, aten.cat, aten.add]
319
+ triton_poi_fused_add_cat_mul_0_xnumel = s0*s1*s7*s8
320
+ stream0 = get_raw_stream(0)
321
+ triton_poi_fused_add_cat_mul_0.run(primals_8, primals_3, primals_5, buf0, s1, ps0, s0, triton_poi_fused_add_cat_mul_0_xnumel, stream=stream0)
322
+ del primals_8
323
+ ps1 = s1*s14
324
+ buf1 = empty_strided_cuda((s13, s14, s0, s1), (s0*s1*s14, s1, s1*s14, 1), torch.bfloat16)
325
+ # Topologically Sorted Source Nodes: [mul_2, cat_1, mul_3, k_embed], Original ATen: [aten.mul, aten.cat, aten.add]
326
+ triton_poi_fused_add_cat_mul_1_xnumel = s0*s1*s13*s14
327
+ stream0 = get_raw_stream(0)
328
+ triton_poi_fused_add_cat_mul_1.run(primals_11, primals_3, primals_5, buf1, s1, ps1, s0, triton_poi_fused_add_cat_mul_1_xnumel, stream=stream0)
329
+ del primals_11
330
+ return (buf0, buf1, primals_3, primals_5, s0, s1, s7, s8, s13, s14, s1 // 2, s1 + (-1)*(s1 // 2), )
331
+
332
+
333
+ def benchmark_compiled_module(times=10, repeat=10):
334
+ from torch._dynamo.testing import rand_strided
335
+ from torch._inductor.utils import print_performance
336
+ primals_1 = 177
337
+ primals_2 = 128
338
+ primals_3 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
339
+ primals_4 = 1
340
+ primals_5 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
341
+ primals_6 = 4
342
+ primals_7 = 16
343
+ primals_8 = rand_strided((4, 16, 177, 128), (362496, 128, 2048, 1), device='cuda:0', dtype=torch.bfloat16)
344
+ primals_9 = 4
345
+ primals_10 = 8
346
+ primals_11 = rand_strided((4, 8, 177, 128), (181248, 128, 1024, 1), device='cuda:0', dtype=torch.bfloat16)
347
+ fn = lambda: call([primals_1, primals_2, primals_3, primals_4, primals_5, primals_6, primals_7, primals_8, primals_9, primals_10, primals_11])
348
+ return print_performance(fn, times=times, repeat=repeat)
349
+
350
+
351
+ if __name__ == "__main__":
352
+ from torch._inductor.wrapper_benchmark import compiled_module_main
353
+ compiled_module_main('None', benchmark_compiled_module)
torchinductor_ch-epfl-345354-j/sw/cswtunku7iwygay5azmeyjzubvchodclgrov3ytwoqfke3cbjeek.py ADDED
@@ -0,0 +1,357 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Compile-time auto-tuning block:
3
+
4
+ import torch
5
+ from torch._dynamo.testing import rand_strided
6
+ from torch._dynamo.utils import preserve_rng_state
7
+ from torch._inductor.select_algorithm import AlgorithmSelectorCache
8
+ from torch._inductor.async_compile import AsyncCompile
9
+
10
+ async_compile = AsyncCompile()
11
+ generate_example_value = AlgorithmSelectorCache.generate_example_value
12
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
13
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
14
+
15
+
16
+ triton_poi_fused_add_cat_mul_0 = async_compile.triton('triton_poi_fused_add_cat_mul_0', '''
17
+ import triton
18
+ import triton.language as tl
19
+
20
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
21
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
22
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
23
+ triton_helpers.set_driver_to_gpu()
24
+
25
+ @triton_heuristics.pointwise(
26
+ size_hints={'x': 2097152},
27
+ filename=__file__,
28
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
29
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
30
+ min_elem_per_thread=0
31
+ )
32
+ @triton.jit
33
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
34
+ xoffset = tl.program_id(0) * XBLOCK
35
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
36
+ xmask = xindex < xnumel
37
+ x4 = xindex
38
+ x0 = (xindex % ks0)
39
+ x2 = ((xindex // ks1) % ks2)
40
+ x5 = xindex // ks0
41
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
42
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
43
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
44
+ tmp2 = tmp0 * tmp1
45
+ tmp3 = x0
46
+ tmp4 = tl.full([1], 0, tl.int64)
47
+ tmp5 = tmp3 >= tmp4
48
+ tmp6 = ks0 + (-1)*(ks0 // 2)
49
+ tmp7 = tmp3 < tmp6
50
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
51
+ tmp9 = -tmp8
52
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
53
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
54
+ tmp12 = tmp3 >= tmp6
55
+ tmp13 = ks0
56
+ tmp14 = tmp3 < tmp13
57
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
58
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
59
+ tmp18 = tmp16 * tmp17
60
+ tmp19 = tmp2 + tmp18
61
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
62
+ ''', device_str='cuda')
63
+
64
+
65
+ triton_poi_fused_add_cat_mul_1 = async_compile.triton('triton_poi_fused_add_cat_mul_1', '''
66
+ import triton
67
+ import triton.language as tl
68
+
69
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
70
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
71
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
72
+ triton_helpers.set_driver_to_gpu()
73
+
74
+ @triton_heuristics.pointwise(
75
+ size_hints={'x': 1048576},
76
+ filename=__file__,
77
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
78
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
79
+ min_elem_per_thread=0
80
+ )
81
+ @triton.jit
82
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
83
+ xoffset = tl.program_id(0) * XBLOCK
84
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
85
+ xmask = xindex < xnumel
86
+ x4 = xindex
87
+ x0 = (xindex % ks0)
88
+ x2 = ((xindex // ks1) % ks2)
89
+ x5 = xindex // ks0
90
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
91
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
92
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
93
+ tmp2 = tmp0 * tmp1
94
+ tmp3 = x0
95
+ tmp4 = tl.full([1], 0, tl.int64)
96
+ tmp5 = tmp3 >= tmp4
97
+ tmp6 = ks0 + (-1)*(ks0 // 2)
98
+ tmp7 = tmp3 < tmp6
99
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
100
+ tmp9 = -tmp8
101
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
102
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
103
+ tmp12 = tmp3 >= tmp6
104
+ tmp13 = ks0
105
+ tmp14 = tmp3 < tmp13
106
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
107
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
108
+ tmp18 = tmp16 * tmp17
109
+ tmp19 = tmp2 + tmp18
110
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
111
+ ''', device_str='cuda')
112
+
113
+ async_compile.wait(globals())
114
+ del async_compile
115
+
116
+ import triton
117
+ import triton.language as tl
118
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
119
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
120
+ with torch.cuda._DeviceGuard(0):
121
+ torch.cuda.set_device(0)
122
+ stream0 = get_raw_stream(0)
123
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
124
+ stream0 = get_raw_stream(0)
125
+ arg7_1 = generate_example_value((4, 16, 177, 128), (362496, 128, 2048, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
126
+ arg2_1 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
127
+ arg4_1 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
128
+ buf0 = generate_example_value((4, 16, 177, 128), (362496, 128, 2048, 1), 'cuda:0', torch.bfloat16, 0, (4, 16, 177, 128))
129
+ triton_poi_fused_add_cat_mul_0.run(arg7_1, arg2_1, arg4_1, buf0, 128, 2048, 177, 1449984, stream=stream0)
130
+ del arg7_1, arg2_1, arg4_1, buf0
131
+
132
+ stream0 = get_raw_stream(0)
133
+ arg10_1 = generate_example_value((4, 8, 177, 128), (181248, 128, 1024, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
134
+ arg2_1 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
135
+ arg4_1 = generate_example_value((1, 177, 128), (22656, 128, 1), 'cuda:0', torch.bfloat16, 0, (1, 177, 128))
136
+ buf1 = generate_example_value((4, 8, 177, 128), (181248, 128, 1024, 1), 'cuda:0', torch.bfloat16, 0, (4, 8, 177, 128))
137
+ triton_poi_fused_add_cat_mul_1.run(arg10_1, arg2_1, arg4_1, buf1, 128, 1024, 177, 724992, stream=stream0)
138
+ del arg10_1, arg2_1, arg4_1, buf1
139
+
140
+ """
141
+ # AOT ID: ['12_inference']
142
+ from ctypes import c_void_p, c_long, c_int
143
+ import torch
144
+ import math
145
+ import random
146
+ import os
147
+ import tempfile
148
+ from math import inf, nan
149
+ from cmath import nanj
150
+ from torch._inductor.hooks import run_intermediate_hooks
151
+ from torch._inductor.utils import maybe_profile
152
+ from torch._inductor.codegen.memory_planning import _align as align
153
+ from torch import device, empty_strided
154
+ from torch._inductor.async_compile import AsyncCompile
155
+ from torch._inductor.select_algorithm import extern_kernels
156
+ from torch._inductor.codegen.multi_kernel import MultiKernelCall
157
+ import triton
158
+ import triton.language as tl
159
+ from torch._inductor.runtime.triton_heuristics import start_graph, end_graph
160
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
161
+ from torch._C import _cuda_getCurrentRawStream as get_raw_stream
162
+
163
+ aten = torch.ops.aten
164
+ inductor_ops = torch.ops.inductor
165
+ _quantized = torch.ops._quantized
166
+ assert_size_stride = torch._C._dynamo.guards.assert_size_stride
167
+ empty_strided_cpu = torch._C._dynamo.guards._empty_strided_cpu
168
+ empty_strided_cuda = torch._C._dynamo.guards._empty_strided_cuda
169
+ empty_strided_xpu = torch._C._dynamo.guards._empty_strided_xpu
170
+ reinterpret_tensor = torch._C._dynamo.guards._reinterpret_tensor
171
+ alloc_from_pool = torch.ops.inductor._alloc_from_pool
172
+ async_compile = AsyncCompile()
173
+ empty_strided_p2p = torch._C._distributed_c10d._SymmetricMemory.empty_strided_p2p
174
+
175
+
176
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/sy/csyofljuv4nybcwnbyfhmdfufvm6xfyypzxkd745lbsexq35kdpk.py
177
+ # Topologically Sorted Source Nodes: [mul, cat, mul_1, q_embed], Original ATen: [aten.mul, aten.cat, aten.add]
178
+ # Source node to ATen node mapping:
179
+ # cat => cat
180
+ # mul => mul_8
181
+ # mul_1 => mul_29
182
+ # q_embed => add_36
183
+ # Graph fragment:
184
+ # %mul_8 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%arg7_1, %unsqueeze), kwargs = {})
185
+ # %cat : [num_users=1] = call_function[target=torch.ops.aten.cat.default](args = ([%neg, %slice_1], -1), kwargs = {})
186
+ # %mul_29 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%cat, %unsqueeze_1), kwargs = {})
187
+ # %add_36 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%mul_8, %mul_29), kwargs = {})
188
+ triton_poi_fused_add_cat_mul_0 = async_compile.triton('triton_poi_fused_add_cat_mul_0', '''
189
+ import triton
190
+ import triton.language as tl
191
+
192
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
193
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
194
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
195
+ triton_helpers.set_driver_to_gpu()
196
+
197
+ @triton_heuristics.pointwise(
198
+ size_hints={'x': 2097152},
199
+ filename=__file__,
200
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
201
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
202
+ min_elem_per_thread=0
203
+ )
204
+ @triton.jit
205
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
206
+ xoffset = tl.program_id(0) * XBLOCK
207
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
208
+ xmask = xindex < xnumel
209
+ x4 = xindex
210
+ x0 = (xindex % ks0)
211
+ x2 = ((xindex // ks1) % ks2)
212
+ x5 = xindex // ks0
213
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
214
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
215
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
216
+ tmp2 = tmp0 * tmp1
217
+ tmp3 = x0
218
+ tmp4 = tl.full([1], 0, tl.int64)
219
+ tmp5 = tmp3 >= tmp4
220
+ tmp6 = ks0 + (-1)*(ks0 // 2)
221
+ tmp7 = tmp3 < tmp6
222
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
223
+ tmp9 = -tmp8
224
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
225
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
226
+ tmp12 = tmp3 >= tmp6
227
+ tmp13 = ks0
228
+ tmp14 = tmp3 < tmp13
229
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
230
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
231
+ tmp18 = tmp16 * tmp17
232
+ tmp19 = tmp2 + tmp18
233
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
234
+ ''', device_str='cuda')
235
+
236
+
237
+ # kernel path: /tmp/torchinductor_ch-epfl-345354-j/da/cdapkyn6r3i3fmlikyctcbjr336rmrrljtxcaqcwrc7auok6odb4.py
238
+ # Topologically Sorted Source Nodes: [mul_2, cat_1, mul_3, k_embed], Original ATen: [aten.mul, aten.cat, aten.add]
239
+ # Source node to ATen node mapping:
240
+ # cat_1 => cat_1
241
+ # k_embed => add_72
242
+ # mul_2 => mul_38
243
+ # mul_3 => mul_59
244
+ # Graph fragment:
245
+ # %mul_38 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%arg10_1, %unsqueeze), kwargs = {})
246
+ # %cat_1 : [num_users=1] = call_function[target=torch.ops.aten.cat.default](args = ([%neg_1, %slice_3], -1), kwargs = {})
247
+ # %mul_59 : [num_users=1] = call_function[target=torch.ops.aten.mul.Tensor](args = (%cat_1, %unsqueeze_1), kwargs = {})
248
+ # %add_72 : [num_users=1] = call_function[target=torch.ops.aten.add.Tensor](args = (%mul_38, %mul_59), kwargs = {})
249
+ triton_poi_fused_add_cat_mul_1 = async_compile.triton('triton_poi_fused_add_cat_mul_1', '''
250
+ import triton
251
+ import triton.language as tl
252
+
253
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
254
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
255
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
256
+ triton_helpers.set_driver_to_gpu()
257
+
258
+ @triton_heuristics.pointwise(
259
+ size_hints={'x': 1048576},
260
+ filename=__file__,
261
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
262
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_1', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
263
+ min_elem_per_thread=0
264
+ )
265
+ @triton.jit
266
+ def triton_poi_fused_add_cat_mul_1(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
267
+ xoffset = tl.program_id(0) * XBLOCK
268
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
269
+ xmask = xindex < xnumel
270
+ x4 = xindex
271
+ x0 = (xindex % ks0)
272
+ x2 = ((xindex // ks1) % ks2)
273
+ x5 = xindex // ks0
274
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
275
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
276
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
277
+ tmp2 = tmp0 * tmp1
278
+ tmp3 = x0
279
+ tmp4 = tl.full([1], 0, tl.int64)
280
+ tmp5 = tmp3 >= tmp4
281
+ tmp6 = ks0 + (-1)*(ks0 // 2)
282
+ tmp7 = tmp3 < tmp6
283
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
284
+ tmp9 = -tmp8
285
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
286
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
287
+ tmp12 = tmp3 >= tmp6
288
+ tmp13 = ks0
289
+ tmp14 = tmp3 < tmp13
290
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
291
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
292
+ tmp18 = tmp16 * tmp17
293
+ tmp19 = tmp2 + tmp18
294
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
295
+ ''', device_str='cuda')
296
+
297
+
298
+ async_compile.wait(globals())
299
+ del async_compile
300
+
301
+ def call(args):
302
+ arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1, arg10_1 = args
303
+ args.clear()
304
+ s0 = arg0_1
305
+ s1 = arg1_1
306
+ s7 = arg5_1
307
+ s8 = arg6_1
308
+ s13 = arg8_1
309
+ s14 = arg9_1
310
+ assert_size_stride(arg2_1, (1, s0, s1), (s0*s1, s1, 1))
311
+ assert_size_stride(arg4_1, (1, s0, s1), (s0*s1, s1, 1))
312
+ assert_size_stride(arg7_1, (s7, s8, s0, s1), (s0*s1*s8, s1, s1*s8, 1))
313
+ assert_size_stride(arg10_1, (s13, s14, s0, s1), (s0*s1*s14, s1, s1*s14, 1))
314
+ with torch.cuda._DeviceGuard(0):
315
+ torch.cuda.set_device(0)
316
+ ps0 = s1*s8
317
+ pool1 = empty_strided_cuda((s7, s8, s0, s1), (s0*s1*s8, s1, s1*s8, 1), torch.bfloat16)
318
+ buf0 = pool1 # alloc
319
+ # Topologically Sorted Source Nodes: [mul, cat, mul_1, q_embed], Original ATen: [aten.mul, aten.cat, aten.add]
320
+ triton_poi_fused_add_cat_mul_0_xnumel = s0*s1*s7*s8
321
+ stream0 = get_raw_stream(0)
322
+ triton_poi_fused_add_cat_mul_0.run(arg7_1, arg2_1, arg4_1, buf0, s1, ps0, s0, triton_poi_fused_add_cat_mul_0_xnumel, stream=stream0)
323
+ del arg7_1
324
+ ps1 = s1*s14
325
+ pool0 = empty_strided_cuda((s13, s14, s0, s1), (s0*s1*s14, s1, s1*s14, 1), torch.bfloat16)
326
+ buf1 = pool0 # alloc
327
+ # Topologically Sorted Source Nodes: [mul_2, cat_1, mul_3, k_embed], Original ATen: [aten.mul, aten.cat, aten.add]
328
+ triton_poi_fused_add_cat_mul_1_xnumel = s0*s1*s13*s14
329
+ stream0 = get_raw_stream(0)
330
+ triton_poi_fused_add_cat_mul_1.run(arg10_1, arg2_1, arg4_1, buf1, s1, ps1, s0, triton_poi_fused_add_cat_mul_1_xnumel, stream=stream0)
331
+ del arg10_1
332
+ del arg2_1
333
+ del arg4_1
334
+ return (buf0, buf1, )
335
+
336
+
337
+ def benchmark_compiled_module(times=10, repeat=10):
338
+ from torch._dynamo.testing import rand_strided
339
+ from torch._inductor.utils import print_performance
340
+ arg0_1 = 177
341
+ arg1_1 = 128
342
+ arg2_1 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
343
+ arg3_1 = 1
344
+ arg4_1 = rand_strided((1, 177, 128), (22656, 128, 1), device='cuda:0', dtype=torch.bfloat16)
345
+ arg5_1 = 4
346
+ arg6_1 = 16
347
+ arg7_1 = rand_strided((4, 16, 177, 128), (362496, 128, 2048, 1), device='cuda:0', dtype=torch.bfloat16)
348
+ arg8_1 = 4
349
+ arg9_1 = 8
350
+ arg10_1 = rand_strided((4, 8, 177, 128), (181248, 128, 1024, 1), device='cuda:0', dtype=torch.bfloat16)
351
+ fn = lambda: call([arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1, arg10_1])
352
+ return print_performance(fn, times=times, repeat=repeat)
353
+
354
+
355
+ if __name__ == "__main__":
356
+ from torch._inductor.wrapper_benchmark import compiled_module_main
357
+ compiled_module_main('None', benchmark_compiled_module)
torchinductor_ch-epfl-345354-j/sy/4d7e479de298eba85b0b0b838eb227f9e0129c56097eee3ae7375c0bd11c1c7d.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 1024, "num_warps": 4, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 25}
torchinductor_ch-epfl-345354-j/sy/csyofljuv4nybcwnbyfhmdfufvm6xfyypzxkd745lbsexq35kdpk.py ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 2097152},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': True, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x4 = xindex
23
+ x0 = (xindex % ks0)
24
+ x2 = ((xindex // ks1) % ks2)
25
+ x5 = xindex // ks0
26
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
28
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
29
+ tmp2 = tmp0 * tmp1
30
+ tmp3 = x0
31
+ tmp4 = tl.full([1], 0, tl.int64)
32
+ tmp5 = tmp3 >= tmp4
33
+ tmp6 = ks0 + (-1)*(ks0 // 2)
34
+ tmp7 = tmp3 < tmp6
35
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
36
+ tmp9 = -tmp8
37
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
38
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
39
+ tmp12 = tmp3 >= tmp6
40
+ tmp13 = ks0
41
+ tmp14 = tmp3 < tmp13
42
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
43
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
44
+ tmp18 = tmp16 * tmp17
45
+ tmp19 = tmp2 + tmp18
46
+ tl.store(out_ptr0 + (x4), tmp19, xmask)
torchinductor_ch-epfl-345354-j/ua/5c2b35108d0f6d838f30c1b016c317aa650f5f275110f17a9aff57ce54e364e6.best_config ADDED
@@ -0,0 +1 @@
 
 
1
+ {"XBLOCK": 1024, "num_warps": 4, "num_stages": 1, "configs_hash": "3ca5c3e34d35093f3c9ab2829a9faeebad5e61c4ca13d5ed6053d7b71ce60d5a", "found_by_coordesc": false, "time_taken_ms": 38}
torchinductor_ch-epfl-345354-j/ua/cuaufinvbemnicqkymrm4ujjq5b6ltwjbdc4jcgwj2zxj4usa5kx.py ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import triton
3
+ import triton.language as tl
4
+
5
+ from torch._inductor.runtime import triton_helpers, triton_heuristics
6
+ from torch._inductor.runtime.triton_helpers import libdevice, math as tl_math
7
+ from torch._inductor.runtime.hints import AutotuneHint, ReductionHint, TileHint, DeviceProperties
8
+ triton_helpers.set_driver_to_gpu()
9
+
10
+ @triton_heuristics.pointwise(
11
+ size_hints={'x': 2097152},
12
+ filename=__file__,
13
+ triton_meta={'signature': {'in_ptr0': '*bf16', 'in_ptr1': '*bf16', 'in_ptr2': '*bf16', 'out_ptr0': '*bf16', 'ks0': 'i32', 'ks1': 'i32', 'ks2': 'i32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='cuda', index=0, multi_processor_count=42, cc=80, major=8, regs_per_multiprocessor=65536, max_threads_per_multi_processor=2048, warp_size=32), 'constants': {}, 'configs': [{(0,): [['tt.divisibility', 16]], (1,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]], (3,): [['tt.divisibility', 16]]}]},
14
+ inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_add_cat_mul_0', 'mutated_arg_names': [], 'optimize_mem': False, 'no_x_dim': False, 'num_load': 5, 'num_reduction': 0, 'backend_hash': '8D9A40F96256AE993B0CB3DAC1136935BA540F7848683690590C84AF795CC5ED', 'are_deterministic_algorithms_enabled': False, 'assert_indirect_indexing': True, 'autotune_local_cache': True, 'autotune_pointwise': True, 'autotune_remote_cache': None, 'force_disable_caches': False, 'dynamic_scale_rblock': True, 'max_autotune': False, 'max_autotune_pointwise': False, 'min_split_scan_rblock': 256, 'spill_threshold': 16, 'store_cubin': False},
15
+ min_elem_per_thread=0
16
+ )
17
+ @triton.jit
18
+ def triton_poi_fused_add_cat_mul_0(in_ptr0, in_ptr1, in_ptr2, out_ptr0, ks0, ks1, ks2, xnumel, XBLOCK : tl.constexpr):
19
+ xoffset = tl.program_id(0) * XBLOCK
20
+ xindex = xoffset + tl.arange(0, XBLOCK)[:]
21
+ xmask = xindex < xnumel
22
+ x4 = xindex
23
+ x0 = (xindex % ks0)
24
+ x2 = ((xindex // ks1) % ks2)
25
+ x5 = xindex // ks0
26
+ tmp0 = tl.load(in_ptr0 + (x4), xmask, eviction_policy='evict_last').to(tl.float32)
27
+ tmp1 = tl.load(in_ptr1 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
28
+ tmp17 = tl.load(in_ptr2 + (x0 + ks0*x2), xmask, eviction_policy='evict_last').to(tl.float32)
29
+ tmp2 = tmp0 * tmp1
30
+ tmp3 = x0
31
+ tmp4 = tl.full([1], 0, tl.int64)
32
+ tmp5 = tmp3 >= tmp4
33
+ tmp6 = ks0 + (-1)*(ks0 // 2)
34
+ tmp7 = tmp3 < tmp6
35
+ tmp8 = tl.load(in_ptr0 + (ks0*x5 + (ks0 // 2) + (x0)), xmask & tmp7, eviction_policy='evict_last', other=0.0).to(tl.float32)
36
+ tmp9 = -tmp8
37
+ tmp10 = tl.full(tmp9.shape, 0.0, tmp9.dtype)
38
+ tmp11 = tl.where(tmp7, tmp9, tmp10)
39
+ tmp12 = tmp3 >= tmp6
40
+ tmp13 = ks0
41
+ tmp14 = tmp3 < tmp13
42
+ tmp15 = tl.load(in_ptr0 + (ks0*x5 + (x0 + ((-1)*ks0) + (ks0 // 2))), xmask & tmp12, eviction_policy='evict_last', other=0.0).to(tl.float32)
43
+ tmp16 = tl.where(tmp7, tmp11, tmp15)
44
+ tmp18 = tmp16 * tmp17
45
+ tmp19 = tmp2 + tmp18
46
+ tl.store(out_ptr0 + (x4), tmp19, xmask)