This is an automated email from the ASF dual-hosted git repository.
junrushao pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/unity by this push:
new 3f1469ea82 [Unity][DLight] Update GEMV rules (#15429)
3f1469ea82 is described below
commit 3f1469ea820901dfb5055e23da7a5a440fd4b1cf
Author: Siyuan Feng <[email protected]>
AuthorDate: Sat Jul 29 16:22:40 2023 +0800
[Unity][DLight] Update GEMV rules (#15429)
This PR updates the GEMV rules:
1. improve for large workloads, speedup Vicuna 13B from 43.0 tok/s to 60.8
tok/s
2. Fix the issue of unexpected unroll for allreduce parts.
---
python/tvm/dlight/gpu/gemv.py | 36 +++---
tests/python/dlight/test_gpu_gemv.py | 215 ++++++++++++++++++-----------------
2 files changed, 132 insertions(+), 119 deletions(-)
diff --git a/python/tvm/dlight/gpu/gemv.py b/python/tvm/dlight/gpu/gemv.py
index 66dac6cb74..0d0e4845d4 100644
--- a/python/tvm/dlight/gpu/gemv.py
+++ b/python/tvm/dlight/gpu/gemv.py
@@ -196,19 +196,28 @@ class GEMV(ScheduleRule):
):
"""Schedule the inner reduction block."""
# pylint: disable=invalid-name
- _, bx, r, _ = sch.get_loops(block)
+ _, s, r, _ = sch.get_loops(block)
# TODO: make it tunable
- len_tx = 32
- len_ty = 8
vec_bytes = 16 if target.kind.name == "cuda" else 8
unroll_number = 256 if target.kind.name == "cuda" else 64
- # Specify the `len_ty` according to the loop extent
- bx_loop: tir.For = sch.get(bx)
- if isinstance(bx_loop.extent, tir.IntImm):
- len_ty = min(bx_loop.extent.value, len_ty)
+ def get_extent(loop_rv: tir.schedule.LoopRV):
+ loop: tir.For = sch.get(loop_rv)
+ return loop.extent.value if isinstance(loop.extent, tir.IntImm)
else 1
+
+ # Specify the `len_tx` and `len_ty` according to the loop extent
+ len_s, len_r = get_extent(s), get_extent(r)
+ if len_r >= 4096 and len_r % 128 == 0:
+ len_tx = 128
+ elif 1024 < len_r <= 2048 and len_r % 64 == 0:
+ len_tx = 64
+ else:
+ len_tx = 32
+
+ if len_s >= 4096:
+ len_ty = 8
else:
- len_ty = 1
+ len_ty = min(len_s, 4)
_, tx = sch.split(r, [None, len_tx], preserve_unit_iters=True)
# Schedule the RF block
@@ -220,8 +229,9 @@ class GEMV(ScheduleRule):
sch.bind(bx, "blockIdx.x")
sch.bind(ty, "threadIdx.y")
sch.bind(tx, "threadIdx.x")
- sch.annotate(tx, "pragma_auto_unroll_max_step", unroll_number)
- sch.annotate(tx, "pragma_unroll_explicit", 1)
+ unit = sch.add_unit_loop(r)
+ sch.annotate(unit, "pragma_auto_unroll_max_step", unroll_number)
+ sch.annotate(unit, "pragma_unroll_explicit", 1)
if target.kind.name == "cuda":
# Cache read the vector
@@ -229,8 +239,8 @@ class GEMV(ScheduleRule):
block: tir.Block = sch.get(rf)
type_bytes: int = get_bytes(block.reads[index].buffer.dtype)
cache = sch.cache_read(rf, index, "shared")
- sch.compute_at(cache, tx, preserve_unit_loops=True)
- fused = sch.fuse(*sch.get_loops(cache)[4:])
+ sch.compute_at(cache, unit, preserve_unit_loops=True)
+ fused = sch.fuse(*sch.get_loops(cache)[5:])
_, _ty, _tx, _vec = sch.split(
fused, [None, len_ty, len_tx, vec_bytes // type_bytes]
)
@@ -244,7 +254,7 @@ class GEMV(ScheduleRule):
vec_length = vec_bytes // type_bytes
cache = sch.cache_read(rf, index, "local")
sch.compute_at(cache, r, preserve_unit_loops=True)
- fused = sch.fuse(*sch.get_loops(cache)[5:])
+ fused = sch.fuse(*sch.get_loops(cache)[6:])
loop: tir.For = sch.get(fused)
if isinstance(loop.extent, tir.IntImm) and loop.extent.value %
vec_length == 0:
_, _vec = sch.split(fused, [None, vec_length])
diff --git a/tests/python/dlight/test_gpu_gemv.py
b/tests/python/dlight/test_gpu_gemv.py
index 1e4245a571..82648fb867 100644
--- a/tests/python/dlight/test_gpu_gemv.py
+++ b/tests/python/dlight/test_gpu_gemv.py
@@ -96,44 +96,45 @@ class TestGEMV(BaseBeforeAfter):
for ax0_fused in T.thread_binding(32, thread="blockIdx.y"):
for ax1_fused_0 in T.thread_binding(n, thread="blockIdx.x"):
for ax1_fused_1 in T.thread_binding(1, thread="threadIdx.y"):
- for ax2_fused_1 in T.thread_binding(32,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
- for ax0_ax1_ax2_ax3_fused_0 in range(1):
- for ax0_ax1_ax2_ax3_fused_1 in T.thread_binding(1,
thread="threadIdx.y"):
- for ax0_ax1_ax2_ax3_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
- for ax0_ax1_ax2_ax3_fused_3 in
T.vectorized(8):
- with T.block("lv1637_shared"):
- v0 = T.axis.spatial(1, 0)
- v1 = T.axis.spatial(32, ax0_fused)
- v2 = T.axis.spatial(1, 0)
- v3 = T.axis.spatial(128,
ax0_ax1_ax2_ax3_fused_0 * 256 + ax0_ax1_ax2_ax3_fused_1 * 256 +
ax0_ax1_ax2_ax3_fused_2 * 8 + ax0_ax1_ax2_ax3_fused_3)
- T.where(((ax0_ax1_ax2_ax3_fused_0
+ ax0_ax1_ax2_ax3_fused_1) * 32 + ax0_ax1_ax2_ax3_fused_2) * 8 +
ax0_ax1_ax2_ax3_fused_3 < 128)
- T.reads(lv1637[v0, v1, v2, v3])
- T.writes(lv1637_shared[v0, v1, v2,
v3])
- lv1637_shared[v0, v1, v2, v3] =
lv1637[v0, v1, v2, v3]
- with T.block("NT_matmul_rf_init"):
- vax2_fused_1, v0 = T.axis.remap("SS",
[ax2_fused_1, ax0_fused])
- v1 = T.axis.spatial(n, ax1_fused_0 + ax1_fused_1)
- T.reads()
-
T.writes(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1])
- var_NT_matmul_intermediate_rf_local[vax2_fused_1,
0, v0, 0, v1] = T.float16(0)
- for ax2_fused_0 in range(4):
- for ax0_ax1_ax2_ax3_fused in T.vectorized(1):
- with T.block("lv1637_shared_local"):
- v0 = T.axis.spatial(1, 0)
- v1 = T.axis.spatial(32, ax0_fused)
- v2 = T.axis.spatial(1, 0)
- v3 = T.axis.spatial(128, ax2_fused_0 * 32
+ ax2_fused_1)
- T.reads(lv1637_shared[v0, v1, v2, v3])
- T.writes(lv1637_shared_local[v0, v1, v2,
v3])
- lv1637_shared_local[v0, v1, v2, v3] =
lv1637_shared[v0, v1, v2, v3]
- for u in range(1):
- with T.block("NT_matmul_rf_update"):
- vax2_fused_1, v0 = T.axis.remap("SS",
[ax2_fused_1, ax0_fused])
- v1 = T.axis.spatial(n, ax1_fused_0 +
ax1_fused_1)
- vax2_fused_0 = T.axis.reduce(4,
ax2_fused_0)
-
T.reads(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1],
lv1637_shared_local[0, v0, 0, vax2_fused_0 * 32 + vax2_fused_1], lv1638[0, v0,
v1, vax2_fused_0 * 32 + vax2_fused_1])
-
T.writes(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1])
-
var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1] =
var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1] +
lv1637_shared_local[0, v0, 0, vax2_fused_0 * 32 + vax2_fused_1] * lv1638[0, v0,
v1, vax2_fused_0 * 32 + vax2_fused_1]
+ for ax2_fused_1 in T.thread_binding(32,
thread="threadIdx.x"):
+ for u in T.serial(1,
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
+ for ax0_ax1_ax2_ax3_fused_0 in range(1):
+ for ax0_ax1_ax2_ax3_fused_1 in
T.thread_binding(1, thread="threadIdx.y"):
+ for ax0_ax1_ax2_ax3_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
+ for ax0_ax1_ax2_ax3_fused_3 in
T.vectorized(8):
+ with T.block("lv1637_shared"):
+ v0 = T.axis.spatial(1, 0)
+ v1 = T.axis.spatial(32,
ax0_fused)
+ v2 = T.axis.spatial(1, 0)
+ v3 = T.axis.spatial(128,
ax0_ax1_ax2_ax3_fused_0 * 256 + ax0_ax1_ax2_ax3_fused_1 * 256 +
ax0_ax1_ax2_ax3_fused_2 * 8 + ax0_ax1_ax2_ax3_fused_3)
+
T.where(((ax0_ax1_ax2_ax3_fused_0 + ax0_ax1_ax2_ax3_fused_1) * 32 +
ax0_ax1_ax2_ax3_fused_2) * 8 + ax0_ax1_ax2_ax3_fused_3 < 128)
+ T.reads(lv1637[v0, v1, v2, v3])
+ T.writes(lv1637_shared[v0, v1,
v2, v3])
+ lv1637_shared[v0, v1, v2, v3]
= lv1637[v0, v1, v2, v3]
+ with T.block("NT_matmul_rf_init"):
+ vax2_fused_1, v0 = T.axis.remap("SS",
[ax2_fused_1, ax0_fused])
+ v1 = T.axis.spatial(n, ax1_fused_0 +
ax1_fused_1)
+ T.reads()
+
T.writes(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1])
+
var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1] = T.float16(0)
+ for ax2_fused_0 in range(4):
+ for ax0_ax1_ax2_ax3_fused in T.vectorized(1):
+ with T.block("lv1637_shared_local"):
+ v0 = T.axis.spatial(1, 0)
+ v1 = T.axis.spatial(32, ax0_fused)
+ v2 = T.axis.spatial(1, 0)
+ v3 = T.axis.spatial(128, ax2_fused_0 *
32 + ax2_fused_1)
+ T.reads(lv1637_shared[v0, v1, v2, v3])
+ T.writes(lv1637_shared_local[v0, v1,
v2, v3])
+ lv1637_shared_local[v0, v1, v2, v3] =
lv1637_shared[v0, v1, v2, v3]
+ for u_1 in range(1):
+ with T.block("NT_matmul_rf_update"):
+ vax2_fused_1, v0 = T.axis.remap("SS",
[ax2_fused_1, ax0_fused])
+ v1 = T.axis.spatial(n, ax1_fused_0 +
ax1_fused_1)
+ vax2_fused_0 = T.axis.reduce(4,
ax2_fused_0)
+
T.reads(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1],
lv1637_shared_local[0, v0, 0, vax2_fused_0 * 32 + vax2_fused_1], lv1638[0, v0,
v1, vax2_fused_0 * 32 + vax2_fused_1])
+
T.writes(var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1])
+
var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1] =
var_NT_matmul_intermediate_rf_local[vax2_fused_1, 0, v0, 0, v1] +
lv1637_shared_local[0, v0, 0, vax2_fused_0 * 32 + vax2_fused_1] * lv1638[0, v0,
v1, vax2_fused_0 * 32 + vax2_fused_1]
for ax1_ax2_fused in range(1):
for ax0 in T.thread_binding(32, thread="threadIdx.x"):
with T.block("NT_matmul"):
@@ -185,42 +186,43 @@ class TestDecodeGEMV1(BaseBeforeAfter):
for u_fused in T.thread_binding(1, thread="blockIdx.y"):
for ax0_fused_0 in T.thread_binding(2752, thread="blockIdx.x"):
for ax0_fused_1 in T.thread_binding(8, thread="threadIdx.y"):
- for ax1_0_fused_1 in T.thread_binding(32,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
- for ax0_ax1_ax2_fused_0 in range(2):
- for ax0_ax1_ax2_fused_1 in T.thread_binding(8,
thread="threadIdx.y"):
- for ax0_ax1_ax2_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
- for ax0_ax1_ax2_fused_3 in T.vectorized(8):
- with T.block("lv1654_shared"):
+ for ax1_0_fused_1 in T.thread_binding(32,
thread="threadIdx.x"):
+ for u in T.serial(1,
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
+ for ax0_ax1_ax2_fused_0 in range(2):
+ for ax0_ax1_ax2_fused_1 in T.thread_binding(8,
thread="threadIdx.y"):
+ for ax0_ax1_ax2_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
+ for ax0_ax1_ax2_fused_3 in
T.vectorized(8):
+ with T.block("lv1654_shared"):
+ v0 = T.axis.spatial(1, 0)
+ v1 = T.axis.spatial(1, 0)
+ v2 = T.axis.spatial(4096,
ax0_ax1_ax2_fused_0 * 2048 + ax0_ax1_ax2_fused_1 * 256 + ax0_ax1_ax2_fused_2 *
8 + ax0_ax1_ax2_fused_3)
+ T.reads(lv1654[v0, v1, v2])
+ T.writes(lv1654_shared[v0, v1,
v2])
+ lv1654_shared[v0, v1, v2] =
lv1654[v0, v1, v2]
+ with T.block("NT_matmul_rf_init"):
+ vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
+ v0 = T.axis.spatial(22016, ax0_fused_0 * 8 +
ax0_fused_1)
+ T.reads()
+
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
+
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
+ for ax1_0_fused_0 in range(16):
+ for ax0_ax1_ax2_fused_0 in range(1):
+ for ax0_ax1_ax2_fused_1 in T.vectorized(8):
+ with T.block("lv1654_shared_local"):
v0 = T.axis.spatial(1, 0)
v1 = T.axis.spatial(1, 0)
- v2 = T.axis.spatial(4096,
ax0_ax1_ax2_fused_0 * 2048 + ax0_ax1_ax2_fused_1 * 256 + ax0_ax1_ax2_fused_2 *
8 + ax0_ax1_ax2_fused_3)
- T.reads(lv1654[v0, v1, v2])
- T.writes(lv1654_shared[v0, v1, v2])
- lv1654_shared[v0, v1, v2] =
lv1654[v0, v1, v2]
- with T.block("NT_matmul_rf_init"):
- vax1_0_fused_1 = T.axis.spatial(32, ax1_0_fused_1)
- v0 = T.axis.spatial(22016, ax0_fused_0 * 8 +
ax0_fused_1)
- T.reads()
-
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
-
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
- for ax1_0_fused_0 in range(16):
- for ax0_ax1_ax2_fused_0 in range(1):
- for ax0_ax1_ax2_fused_1 in T.vectorized(8):
- with T.block("lv1654_shared_local"):
- v0 = T.axis.spatial(1, 0)
- v1 = T.axis.spatial(1, 0)
- v2 = T.axis.spatial(4096,
ax1_0_fused_0 * 256 + ax1_0_fused_1 * 8 + ax0_ax1_ax2_fused_0 * 8 +
ax0_ax1_ax2_fused_1)
- T.reads(lv1654_shared[v0, v1, v2])
- T.writes(lv1654_shared_local[v0, v1,
v2])
- lv1654_shared_local[v0, v1, v2] =
lv1654_shared[v0, v1, v2]
- for ax1_1 in range(8):
- with T.block("NT_matmul_rf_update"):
- vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
- v0 = T.axis.spatial(22016, ax0_fused_0 * 8
+ ax0_fused_1)
- vax1_0_fused_0, vax1_1 =
T.axis.remap("RR", [ax1_0_fused_0, ax1_1])
-
T.reads(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0],
lv1654_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1],
lv571[v0, (vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 8], lv572[v0,
(vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 32])
-
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
-
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] =
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] +
lv1654_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1] *
((T.Cast("float16", T.bitwise_and(T.shift_right(lv571[v0, (vax1_0_fused_0 * 256
+ vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 256 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * lv572[ [...]
+ v2 = T.axis.spatial(4096,
ax1_0_fused_0 * 256 + ax1_0_fused_1 * 8 + ax0_ax1_ax2_fused_0 * 8 +
ax0_ax1_ax2_fused_1)
+ T.reads(lv1654_shared[v0, v1, v2])
+ T.writes(lv1654_shared_local[v0,
v1, v2])
+ lv1654_shared_local[v0, v1, v2] =
lv1654_shared[v0, v1, v2]
+ for ax1_1 in range(8):
+ with T.block("NT_matmul_rf_update"):
+ vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
+ v0 = T.axis.spatial(22016, ax0_fused_0
* 8 + ax0_fused_1)
+ vax1_0_fused_0, vax1_1 =
T.axis.remap("RR", [ax1_0_fused_0, ax1_1])
+
T.reads(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0],
lv1654_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1],
lv571[v0, (vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 8], lv572[v0,
(vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
+
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] =
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] +
lv1654_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1] *
((T.Cast("float16", T.bitwise_and(T.shift_right(lv571[v0, (vax1_0_fused_0 * 256
+ vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 256 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * lv [...]
for ax1_fused in range(1):
for ax0 in T.thread_binding(32, thread="threadIdx.x"):
with T.block("NT_matmul"):
@@ -276,42 +278,43 @@ class TestDecodeGEMV2(BaseBeforeAfter):
for u_fused in T.thread_binding(1, thread="blockIdx.y"):
for ax0_fused_0 in T.thread_binding(4000, thread="blockIdx.x"):
for ax0_fused_1 in T.thread_binding(8, thread="threadIdx.y"):
- for ax1_0_fused_1 in T.thread_binding(32,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
- for ax0_ax1_ax2_fused_0 in range(2):
- for ax0_ax1_ax2_fused_1 in T.thread_binding(8,
thread="threadIdx.y"):
- for ax0_ax1_ax2_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
- for ax0_ax1_ax2_fused_3 in T.vectorized(8):
- with T.block("lv3216_shared"):
+ for ax1_0_fused_1 in T.thread_binding(32,
thread="threadIdx.x"):
+ for u in T.serial(1,
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
+ for ax0_ax1_ax2_fused_0 in range(2):
+ for ax0_ax1_ax2_fused_1 in T.thread_binding(8,
thread="threadIdx.y"):
+ for ax0_ax1_ax2_fused_2 in
T.thread_binding(32, thread="threadIdx.x"):
+ for ax0_ax1_ax2_fused_3 in
T.vectorized(8):
+ with T.block("lv3216_shared"):
+ v0 = T.axis.spatial(1, 0)
+ v1 = T.axis.spatial(1, 0)
+ v2 = T.axis.spatial(4096,
ax0_ax1_ax2_fused_0 * 2048 + ax0_ax1_ax2_fused_1 * 256 + ax0_ax1_ax2_fused_2 *
8 + ax0_ax1_ax2_fused_3)
+ T.reads(lv3216[v0, v1, v2])
+ T.writes(lv3216_shared[v0, v1,
v2])
+ lv3216_shared[v0, v1, v2] =
lv3216[v0, v1, v2]
+ with T.block("NT_matmul_rf_init"):
+ vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
+ v0 = T.axis.spatial(32000, ax0_fused_0 * 8 +
ax0_fused_1)
+ T.reads()
+
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
+
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
+ for ax1_0_fused_0 in range(16):
+ for ax0_ax1_ax2_fused_0 in range(1):
+ for ax0_ax1_ax2_fused_1 in T.vectorized(8):
+ with T.block("lv3216_shared_local"):
v0 = T.axis.spatial(1, 0)
v1 = T.axis.spatial(1, 0)
- v2 = T.axis.spatial(4096,
ax0_ax1_ax2_fused_0 * 2048 + ax0_ax1_ax2_fused_1 * 256 + ax0_ax1_ax2_fused_2 *
8 + ax0_ax1_ax2_fused_3)
- T.reads(lv3216[v0, v1, v2])
- T.writes(lv3216_shared[v0, v1, v2])
- lv3216_shared[v0, v1, v2] =
lv3216[v0, v1, v2]
- with T.block("NT_matmul_rf_init"):
- vax1_0_fused_1 = T.axis.spatial(32, ax1_0_fused_1)
- v0 = T.axis.spatial(32000, ax0_fused_0 * 8 +
ax0_fused_1)
- T.reads()
-
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
-
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
- for ax1_0_fused_0 in range(16):
- for ax0_ax1_ax2_fused_0 in range(1):
- for ax0_ax1_ax2_fused_1 in T.vectorized(8):
- with T.block("lv3216_shared_local"):
- v0 = T.axis.spatial(1, 0)
- v1 = T.axis.spatial(1, 0)
- v2 = T.axis.spatial(4096,
ax1_0_fused_0 * 256 + ax1_0_fused_1 * 8 + ax0_ax1_ax2_fused_0 * 8 +
ax0_ax1_ax2_fused_1)
- T.reads(lv3216_shared[v0, v1, v2])
- T.writes(lv3216_shared_local[v0, v1,
v2])
- lv3216_shared_local[v0, v1, v2] =
lv3216_shared[v0, v1, v2]
- for ax1_1 in range(8):
- with T.block("NT_matmul_rf_update"):
- vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
- v0 = T.axis.spatial(32000, ax0_fused_0 * 8
+ ax0_fused_1)
- vax1_0_fused_0, vax1_1 =
T.axis.remap("RR", [ax1_0_fused_0, ax1_1])
-
T.reads(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0],
lv3216_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1],
lv771[v0, (vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 8], lv772[v0,
(vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 32])
-
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
-
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] =
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] +
lv3216_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1] *
((T.Cast("float16", T.bitwise_and(T.shift_right(lv771[v0, (vax1_0_fused_0 * 256
+ vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 256 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * lv772[ [...]
+ v2 = T.axis.spatial(4096,
ax1_0_fused_0 * 256 + ax1_0_fused_1 * 8 + ax0_ax1_ax2_fused_0 * 8 +
ax0_ax1_ax2_fused_1)
+ T.reads(lv3216_shared[v0, v1, v2])
+ T.writes(lv3216_shared_local[v0,
v1, v2])
+ lv3216_shared_local[v0, v1, v2] =
lv3216_shared[v0, v1, v2]
+ for ax1_1 in range(8):
+ with T.block("NT_matmul_rf_update"):
+ vax1_0_fused_1 = T.axis.spatial(32,
ax1_0_fused_1)
+ v0 = T.axis.spatial(32000, ax0_fused_0
* 8 + ax0_fused_1)
+ vax1_0_fused_0, vax1_1 =
T.axis.remap("RR", [ax1_0_fused_0, ax1_1])
+
T.reads(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0],
lv3216_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1],
lv771[v0, (vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 8], lv772[v0,
(vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+
T.writes(var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0])
+
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] =
var_NT_matmul_intermediate_rf_local[vax1_0_fused_1, 0, 0, v0] +
lv3216_shared_local[0, 0, vax1_0_fused_0 * 256 + vax1_0_fused_1 * 8 + vax1_1] *
((T.Cast("float16", T.bitwise_and(T.shift_right(lv771[v0, (vax1_0_fused_0 * 256
+ vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 256 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * lv [...]
for ax1_fused in range(1):
for ax0 in T.thread_binding(32, thread="threadIdx.x"):
with T.block("NT_matmul"):