Training in progress, step 680
Browse files- .gitattributes +3 -0
- model.safetensors +1 -1
- torchinductor_ch-epfl-345354-j/aotautograd/akum3jin35di7tlhk4tvtzrbeforlbveooabshfrmtcm4cwst6h3/entry +0 -0
- torchinductor_ch-epfl-345354-j/aotautograd/azlrpmzq25mv3r2ycoghvuwp5ap446dq3pwd3e73xi4lravn37ws/entry +0 -0
- torchinductor_ch-epfl-345354-j/bw/cbw2ylnzyqqvpuutlhdcnot22eonl4ebucgx7hpqh66fjy7zgkhb.py +361 -0
- torchinductor_ch-epfl-345354-j/da/0df06b8658127e5ce915248b1384d7c58762d116d843ac58b2d4be91c249301b.best_config +1 -0
- torchinductor_ch-epfl-345354-j/da/cdapkyn6r3i3fmlikyctcbjr336rmrrljtxcaqcwrc7auok6odb4.py +46 -0
- torchinductor_ch-epfl-345354-j/fj/cfjpgg4ctgzyci7wzylrvpq3i7o2o5r6brzgt5u6bzfhltzrjid3.py +48 -0
- torchinductor_ch-epfl-345354-j/fj/fa90c6e1e9df4537c3b4e7a14a704c2df4c98a82068ce1f3abb21af8a35247a0.best_config +1 -0
- torchinductor_ch-epfl-345354-j/fxgraph/3o/f3odjl3g6dddsqtuut46oufxq75twuknijzu35ozpse6s2ndf7rh/ssax6qk4l2jsc4ovqf64hhenn3rpvkiqhluzs7qfxcc6dzbigdx +3 -0
- torchinductor_ch-epfl-345354-j/fxgraph/7y/f7yf7bpzpliuud5ylt3deq7gx7odarlyvhr2zdhcorkmqjxqlpv7/inrn63eoovhcqnqsled5nyfdmzwunnsaud5rcn37go5egvyuze3 +3 -0
- torchinductor_ch-epfl-345354-j/fxgraph/p3/fp32tea4g6rsybgfroshvlbxseks7iiqpb4cmbb2d3lgralck774/vbrtzv3sqiuroqkjniy4wufikxjgiptmxb4qwrykjuzdrp2mvvt +3 -0
- torchinductor_ch-epfl-345354-j/j6/42b81ae014f6f1c38797163131fc53a645adf2cc0424a9255831be278e34af78.best_config +1 -0
- torchinductor_ch-epfl-345354-j/j6/cj6yl63qj32gvorczxkszjnojwnybujlaozjrnzwzik4ukxq4rss.py +48 -0
- torchinductor_ch-epfl-345354-j/lm/65cc8fd7bae3abb18e0f1b23c8e1c60711fe732716b5724936da88bfeb6409b4.best_config +1 -0
- torchinductor_ch-epfl-345354-j/lm/clmeyn6qnatdy2hjtvp2smmgjoj5zqvy23kyx4cn7bbcafehg7wr.py +46 -0
- torchinductor_ch-epfl-345354-j/mc/cmcqiflrggybh55qgleeqzkbj3rlpmylxaqgl47omwv5z3y4ifv5.py +353 -0
- torchinductor_ch-epfl-345354-j/sw/cswtunku7iwygay5azmeyjzubvchodclgrov3ytwoqfke3cbjeek.py +357 -0
- torchinductor_ch-epfl-345354-j/sy/4d7e479de298eba85b0b0b838eb227f9e0129c56097eee3ae7375c0bd11c1c7d.best_config +1 -0
- torchinductor_ch-epfl-345354-j/sy/csyofljuv4nybcwnbyfhmdfufvm6xfyypzxkd745lbsexq35kdpk.py +46 -0
- torchinductor_ch-epfl-345354-j/ua/5c2b35108d0f6d838f30c1b016c317aa650f5f275110f17a9aff57ce54e364e6.best_config +1 -0
- 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:
|
| 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)
|