This is an automated email from the ASF dual-hosted git repository.
hongyij 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 03fc4f6f03 [Dlight] Change max_threads on CUDA (#16203)
03fc4f6f03 is described below
commit 03fc4f6f0314382be602e596b590c87e7d7d55db
Author: Hongyi Jin <[email protected]>
AuthorDate: Thu Dec 7 21:38:54 2023 -0500
[Dlight] Change max_threads on CUDA (#16203)
* change cuda thread num
* fix test
* fix lint
---
python/tvm/dlight/gpu/utils.py | 2 +-
tests/python/dlight/test_gpu_reduction.py | 343 ++++++++++++++++++------------
2 files changed, 203 insertions(+), 142 deletions(-)
diff --git a/python/tvm/dlight/gpu/utils.py b/python/tvm/dlight/gpu/utils.py
index 00d97ab7f1..4f2df5cfa0 100644
--- a/python/tvm/dlight/gpu/utils.py
+++ b/python/tvm/dlight/gpu/utils.py
@@ -50,7 +50,7 @@ def suggest_threads_per_block(
max_threads_for_dynamic_loop: int = 32,
) -> List[int]:
if target.kind.name == "cuda":
- threads = 256
+ threads = 1024
elif target.kind.name == "rocm":
threads = 256
elif target.kind.name == "metal":
diff --git a/tests/python/dlight/test_gpu_reduction.py
b/tests/python/dlight/test_gpu_reduction.py
index 6198a2eb72..ac34d1ee91 100644
--- a/tests/python/dlight/test_gpu_reduction.py
+++ b/tests/python/dlight/test_gpu_reduction.py
@@ -53,29 +53,44 @@ def test_decode_gemv_1():
@I.ir_module
class After:
@T.prim_func
- def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128),
"float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096),
"float16")):
+ def func(W_handle: T.handle, S_handle: T.handle, V_handle: T.handle,
C_handle: T.handle):
T.func_attr({"global_symbol": "main", "tir.is_scheduled": 1,
"tir.noalias": T.bool(True)})
- # with T.block("root"):
- C_rf_local = T.alloc_buffer((256, 1, 1, 4096), "float16",
scope="local")
- for i2_i0_i1_fused in T.thread_binding(4096, thread="blockIdx.x"):
- for k_0_fused_1 in T.thread_binding(256, thread="threadIdx.x",
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
- with T.block("matmul_rf_init"):
- vk_0_fused_1 = T.axis.spatial(256, k_0_fused_1)
- v_i2 = T.axis.spatial(4096, i2_i0_i1_fused)
- C_rf_local[vk_0_fused_1, 0, 0, v_i2] = T.float16(0)
- for k_0_fused_0, k_1 in T.grid(2, 8):
- with T.block("matmul_rf_update"):
- vk_0_fused_1 = T.axis.spatial(256, k_0_fused_1)
- v_i2, vk_0_fused_0, vk_1 = T.axis.remap("SRR",
[i2_i0_i1_fused, k_0_fused_0, k_1])
- C_rf_local[vk_0_fused_1, 0, 0, v_i2] =
C_rf_local[vk_0_fused_1, 0, 0, v_i2] + V[0, 0, vk_0_fused_0 * 2048 +
vk_0_fused_1 * 8 + vk_1] * ((T.Cast("float16",
T.bitwise_and(T.shift_right(W[v_i2, (vk_0_fused_0 * 2048 + vk_0_fused_1 * 8 +
vk_1) // 8], T.Cast("uint32", (vk_0_fused_0 * 2048 + vk_0_fused_1 * 8 + vk_1) %
8) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * S[v_i2, (vk_0_fused_0 *
2048 + vk_0_fused_1 * 8 + vk_1) // 32])
- for ax1_ax2_ax3_fused in range(1): # pylint:
disable=unused-variable
- for ax0_fused in T.thread_binding(256,
thread="threadIdx.x"):
- with T.block("matmul"):
- vk_0_fused_1 = T.axis.reduce(256, ax0_fused)
- v_i2 = T.axis.spatial(4096, i2_i0_i1_fused)
- with T.init():
- C[0, 0, v_i2] = T.float16(0)
- C[0, 0, v_i2] = C[0, 0, v_i2] +
C_rf_local[vk_0_fused_1, 0, 0, v_i2]
+ W = T.match_buffer(W_handle, (4096, 512), "uint32")
+ S = T.match_buffer(S_handle, (4096, 128), "float16")
+ V = T.match_buffer(V_handle, (1, 1, 4096), "float16")
+ C = T.match_buffer(C_handle, (1, 1, 4096), "float16")
+ with T.block("root"):
+ T.reads()
+ T.writes()
+ C_rf_local = T.alloc_buffer((512, 1, 1, 4096), "float16",
scope="local")
+ for ax0_fused in T.thread_binding(4096, thread="blockIdx.x"):
+ for ax1_0_fused_1 in T.thread_binding(512,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
+ with T.block("matmul_rf_init"):
+ vax1_0_fused_1 = T.axis.spatial(512, ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads()
+ T.writes(C_rf_local[vax1_0_fused_1, 0, 0, v0])
+ C_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
+ for ax1_0_fused_0 in range(1):
+ for ax1_1 in range(8):
+ with T.block("matmul_rf_update"):
+ vax1_0_fused_1 = T.axis.spatial(512,
ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ vax1_0_fused_0 = T.axis.reduce(1,
ax1_0_fused_0)
+ vax1_1 = T.axis.reduce(8, ax1_1)
+ T.reads(C_rf_local[vax1_0_fused_1, 0, 0,
v0], V[0, 0, vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1], W[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 8], S[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ T.writes(C_rf_local[vax1_0_fused_1, 0, 0,
v0])
+ C_rf_local[vax1_0_fused_1, 0, 0, v0] =
C_rf_local[vax1_0_fused_1, 0, 0, v0] + V[0, 0, vax1_0_fused_0 * 4096 +
vax1_0_fused_1 * 8 + vax1_1] * ((T.Cast("float16",
T.bitwise_and(T.shift_right(W[v0, (vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 +
vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 +
vax1_1) % 8) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * S[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ for ax1_fused in range(1):
+ for ax0 in T.thread_binding(512, thread="threadIdx.x"):
+ with T.block("matmul"):
+ vax1_0_fused_1 = T.axis.reduce(512, ax0)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads(C_rf_local[vax1_0_fused_1, 0, 0, v0])
+ T.writes(C[0, 0, v0])
+ with T.init():
+ C[0, 0, v0] = T.float16(0)
+ C[0, 0, v0] = C[0, 0, v0] +
C_rf_local[vax1_0_fused_1, 0, 0, v0]
# fmt: on
target = Target("nvidia/geforce-rtx-3090-ti")
@@ -172,36 +187,48 @@ def test_decode_gemv_3():
C[v_i0, v_i1, v_i2] = T.float16(0)
C[v_i0, v_i1, v_i2] = C[v_i0, v_i1, v_i2] + V[v_i0, v_i1,
v_k] * B[v_i2, v_k]
-
@I.ir_module
class After:
@T.prim_func
- def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096),
"float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096),
"float16")):
+ def func(W_handle: T.handle, S_handle: T.handle, V_handle: T.handle,
C_handle: T.handle):
T.func_attr({"global_symbol": "main", "tir.is_scheduled": 1,
"tir.noalias": T.bool(True)})
- # with T.block("root"):
- C_rf_local = T.alloc_buffer((256, 1, 1, 4096), "float16",
scope="local")
- for i2_0_i0_i1_fused in T.thread_binding(512, thread="blockIdx.x"):
- for k_fused_1 in T.thread_binding(256, thread="threadIdx.x",
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
- for i2_1_init in range(8):
- with T.block("matmul_rf_init"):
- vk_fused_1 = T.axis.spatial(256, k_fused_1)
- v_i2 = T.axis.spatial(4096, i2_0_i0_i1_fused * 8 +
i2_1_init)
- C_rf_local[vk_fused_1, 0, 0, v_i2] = T.float16(0)
- for k_fused_0, i2_1 in T.grid(16, 8):
- with T.block("matmul_rf_update"):
- vk_fused_1 = T.axis.spatial(256, k_fused_1)
- v_i2 = T.axis.spatial(4096, i2_0_i0_i1_fused * 8 +
i2_1)
- vk_fused_0 = T.axis.reduce(16, k_fused_0)
- C_rf_local[vk_fused_1, 0, 0, v_i2] =
C_rf_local[vk_fused_1, 0, 0, v_i2] + V[0, 0, vk_fused_0 * 256 + vk_fused_1] *
((T.Cast("float16", T.bitwise_and(T.shift_right(W[v_i2 // 8, vk_fused_0 * 256 +
vk_fused_1], T.Cast("uint32", v_i2 % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * S[v_i2 // 32, vk_fused_0 * 256 + vk_fused_1])
- for ax1_ax2_ax3_fused_0 in range(1):
- for ax0_fused in T.thread_binding(256,
thread="threadIdx.x"):
- for ax1_ax2_ax3_fused_1 in range(8):
- with T.block("matmul"):
- vk_fused_1 = T.axis.reduce(256, ax0_fused)
- v_i2 = T.axis.spatial(4096, i2_0_i0_i1_fused *
8 + ax1_ax2_ax3_fused_0 * 8 + ax1_ax2_ax3_fused_1)
- with T.init():
- C[0, 0, v_i2] = T.float16(0)
- C[0, 0, v_i2] = C[0, 0, v_i2] +
C_rf_local[vk_fused_1, 0, 0, v_i2]
+ W = T.match_buffer(W_handle, (512, 4096), "uint32")
+ S = T.match_buffer(S_handle, (128, 4096), "float16")
+ V = T.match_buffer(V_handle, (1, 1, 4096), "float16")
+ C = T.match_buffer(C_handle, (1, 1, 4096), "float16")
+ with T.block("root"):
+ T.reads()
+ T.writes()
+ C_rf_local = T.alloc_buffer((1024, 1, 1, 4096), "float16",
scope="local")
+ for ax0_0_fused in T.thread_binding(512, thread="blockIdx.x"):
+ for ax1_fused_1 in T.thread_binding(1024,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
+ for ax0_1_init in range(8):
+ with T.block("matmul_rf_init"):
+ vax1_fused_1 = T.axis.spatial(1024,
ax1_fused_1)
+ v0 = T.axis.spatial(4096, ax0_0_fused * 8 +
ax0_1_init)
+ T.reads()
+ T.writes(C_rf_local[vax1_fused_1, 0, 0, v0])
+ C_rf_local[vax1_fused_1, 0, 0, v0] =
T.float16(0)
+ for ax1_fused_0 in range(4):
+ for ax0_1 in range(8):
+ with T.block("matmul_rf_update"):
+ vax1_fused_1 = T.axis.spatial(1024,
ax1_fused_1)
+ v0 = T.axis.spatial(4096, ax0_0_fused * 8
+ ax0_1)
+ vax1_fused_0 = T.axis.reduce(4,
ax1_fused_0)
+ T.reads(C_rf_local[vax1_fused_1, 0, 0,
v0], V[0, 0, vax1_fused_0 * 1024 + vax1_fused_1], W[v0 // 8, vax1_fused_0 *
1024 + vax1_fused_1], S[v0 // 32, vax1_fused_0 * 1024 + vax1_fused_1])
+ T.writes(C_rf_local[vax1_fused_1, 0, 0,
v0])
+ C_rf_local[vax1_fused_1, 0, 0, v0] =
C_rf_local[vax1_fused_1, 0, 0, v0] + V[0, 0, vax1_fused_0 * 1024 +
vax1_fused_1] * ((T.Cast("float16", T.bitwise_and(T.shift_right(W[v0 // 8,
vax1_fused_0 * 1024 + vax1_fused_1], T.Cast("uint32", v0 % 8) * T.uint32(4)),
T.uint32(15))) - T.float16(7)) * S[v0 // 32, vax1_fused_0 * 1024 +
vax1_fused_1])
+ for ax1_fused_0 in range(1):
+ for ax0 in T.thread_binding(1024,
thread="threadIdx.x"):
+ for ax1_fused_1 in range(8):
+ with T.block("matmul"):
+ vax1_fused_1 = T.axis.reduce(1024, ax0)
+ v0 = T.axis.spatial(4096, ax0_0_fused * 8
+ ax1_fused_0 * 8 + ax1_fused_1)
+ T.reads(C_rf_local[vax1_fused_1, 0, 0, v0])
+ T.writes(C[0, 0, v0])
+ with T.init():
+ C[0, 0, v0] = T.float16(0)
+ C[0, 0, v0] = C[0, 0, v0] +
C_rf_local[vax1_fused_1, 0, 0, v0]
# fmt: on
@@ -311,33 +338,50 @@ def test_decode_gemv_sigmoid():
@I.ir_module
class After:
@T.prim_func
- def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128),
"float16"), V: T.Buffer((1, 1, 4096), "float16"), D: T.Buffer((1, 1, 4096),
"float16")):
+ def func(W_handle: T.handle, S_handle: T.handle, V_handle: T.handle,
D_handle: T.handle):
T.func_attr({"global_symbol": "main", "tir.is_scheduled": 1,
"tir.noalias": T.bool(True)})
- # with T.block("root"):
- C_local = T.alloc_buffer((1, 1, 4096), "float16", scope="local")
- C_rf_local = T.alloc_buffer((256, 1, 1, 4096), "float16",
scope="local")
- for i2_i0_i1_fused in T.thread_binding(4096, thread="blockIdx.x"):
- for k_0_fused_1 in T.thread_binding(256, thread="threadIdx.x",
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
- with T.block("matmul_rf_init"):
- vk_0_fused_1 = T.axis.spatial(256, k_0_fused_1)
- v_i2 = T.axis.spatial(4096, i2_i0_i1_fused)
- C_rf_local[vk_0_fused_1, 0, 0, v_i2] = T.float16(0)
- for k_0_fused_0, k_1 in T.grid(2, 8):
- with T.block("matmul_rf_update"):
- vk_0_fused_1 = T.axis.spatial(256, k_0_fused_1)
- v_i2, vk_0_fused_0, vk_1 = T.axis.remap("SRR",
[i2_i0_i1_fused, k_0_fused_0, k_1])
- C_rf_local[vk_0_fused_1, 0, 0, v_i2] =
C_rf_local[vk_0_fused_1, 0, 0, v_i2] + V[0, 0, vk_0_fused_0 * 2048 +
vk_0_fused_1 * 8 + vk_1] * ((T.Cast("float16",
T.bitwise_and(T.shift_right(W[v_i2, (vk_0_fused_0 * 2048 + vk_0_fused_1 * 8 +
vk_1) // 8], T.Cast("uint32", (vk_0_fused_0 * 2048 + vk_0_fused_1 * 8 + vk_1) %
8) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * S[v_i2, (vk_0_fused_0 *
2048 + vk_0_fused_1 * 8 + vk_1) // 32])
- for ax1_ax2_ax3_fused in range(1): # pylint:
disable=unused-variable
- for ax0_fused in T.thread_binding(256,
thread="threadIdx.x"):
- with T.block("matmul"):
- vk_0_fused_1 = T.axis.reduce(256, ax0_fused)
- v_i2 = T.axis.spatial(4096, i2_i0_i1_fused)
- with T.init():
- C_local[0, 0, v_i2] = T.float16(0)
- C_local[0, 0, v_i2] = C_local[0, 0, v_i2] +
C_rf_local[vk_0_fused_1, 0, 0, v_i2]
- with T.block("sigmoid"):
- v_i2 = T.axis.spatial(4096, i2_i0_i1_fused)
- D[0, 0, v_i2] = T.sigmoid(C_local[0, 0, v_i2])
+ W = T.match_buffer(W_handle, (4096, 512), "uint32")
+ S = T.match_buffer(S_handle, (4096, 128), "float16")
+ V = T.match_buffer(V_handle, (1, 1, 4096), "float16")
+ D = T.match_buffer(D_handle, (1, 1, 4096), "float16")
+ with T.block("root"):
+ T.reads()
+ T.writes()
+ C_local = T.alloc_buffer((1, 1, 4096), "float16",
scope="local")
+ C_rf_local = T.alloc_buffer((512, 1, 1, 4096), "float16",
scope="local")
+ for ax0_fused in T.thread_binding(4096, thread="blockIdx.x"):
+ for ax1_0_fused_1 in T.thread_binding(512,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
+ with T.block("matmul_rf_init"):
+ vax1_0_fused_1 = T.axis.spatial(512, ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads()
+ T.writes(C_rf_local[vax1_0_fused_1, 0, 0, v0])
+ C_rf_local[vax1_0_fused_1, 0, 0, v0] = T.float16(0)
+ for ax1_0_fused_0 in range(1):
+ for ax1_1 in range(8):
+ with T.block("matmul_rf_update"):
+ vax1_0_fused_1 = T.axis.spatial(512,
ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ vax1_0_fused_0 = T.axis.reduce(1,
ax1_0_fused_0)
+ vax1_1 = T.axis.reduce(8, ax1_1)
+ T.reads(C_rf_local[vax1_0_fused_1, 0, 0,
v0], V[0, 0, vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1], W[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 8], S[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ T.writes(C_rf_local[vax1_0_fused_1, 0, 0,
v0])
+ C_rf_local[vax1_0_fused_1, 0, 0, v0] =
C_rf_local[vax1_0_fused_1, 0, 0, v0] + V[0, 0, vax1_0_fused_0 * 4096 +
vax1_0_fused_1 * 8 + vax1_1] * ((T.Cast("float16",
T.bitwise_and(T.shift_right(W[v0, (vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 +
vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 +
vax1_1) % 8) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * S[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ for ax1_fused in range(1):
+ for ax0 in T.thread_binding(512, thread="threadIdx.x"):
+ with T.block("matmul"):
+ vax1_0_fused_1 = T.axis.reduce(512, ax0)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads(C_rf_local[vax1_0_fused_1, 0, 0, v0])
+ T.writes(C_local[0, 0, v0])
+ with T.init():
+ C_local[0, 0, v0] = T.float16(0)
+ C_local[0, 0, v0] = C_local[0, 0, v0] +
C_rf_local[vax1_0_fused_1, 0, 0, v0]
+ with T.block("sigmoid"):
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads(C_local[0, 0, v0])
+ T.writes(D[0, 0, v0])
+ D[0, 0, v0] = T.sigmoid(C_local[0, 0, v0])
# fmt: on
@@ -379,42 +423,53 @@ def test_decode_gemv_1_fp32():
T.writes(C[v_i0, v_i1, v_i2])
C[v_i0, v_i1, v_i2] = T.Cast("float16", C_fp32[v_i0, v_i1,
v_i2])
-
@I.ir_module
class After:
@T.prim_func
- def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128),
"float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096),
"float16")):
+ def func(W_handle: T.handle, S_handle: T.handle, V_handle: T.handle,
C_handle: T.handle):
T.func_attr({"global_symbol": "main", "tir.is_scheduled": 1,
"tir.noalias": T.bool(True)})
- # with T.block("root"):
- C_fp32_local = T.alloc_buffer((1, 1, 4096), scope="local")
- C_fp32_rf_local = T.alloc_buffer((256, 1, 1, 4096), scope="local")
- for ax0_fused in T.thread_binding(4096, thread="blockIdx.x"):
- for ax1_0_fused_1 in T.thread_binding(256,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
- with T.block("matmul_rf_init"):
- vax1_0_fused_1, v0 = T.axis.remap("SS",
[ax1_0_fused_1, ax0_fused])
- T.reads()
- T.writes(C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0])
- C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0] =
T.float32(0)
- for ax1_0_fused_0, ax1_1 in T.grid(2, 8):
- with T.block("matmul_rf_update"):
- vax1_0_fused_1, v0, vax1_0_fused_0, vax1_1 =
T.axis.remap("SSRR", [ax1_0_fused_1, ax0_fused, ax1_0_fused_0, ax1_1])
- T.reads(C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0],
V[0, 0, vax1_0_fused_0 * 2048 + vax1_0_fused_1 * 8 + vax1_1], W[v0,
(vax1_0_fused_0 * 2048 + vax1_0_fused_1 * 8 + vax1_1) // 8], S[v0,
(vax1_0_fused_0 * 2048 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ W = T.match_buffer(W_handle, (4096, 512), "uint32")
+ S = T.match_buffer(S_handle, (4096, 128), "float16")
+ V = T.match_buffer(V_handle, (1, 1, 4096), "float16")
+ C = T.match_buffer(C_handle, (1, 1, 4096), "float16")
+ with T.block("root"):
+ T.reads()
+ T.writes()
+ C_fp32_local = T.alloc_buffer((1, 1, 4096), scope="local")
+ C_fp32_rf_local = T.alloc_buffer((512, 1, 1, 4096),
scope="local")
+ for ax0_fused in T.thread_binding(4096, thread="blockIdx.x"):
+ for ax1_0_fused_1 in T.thread_binding(512,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
+ with T.block("matmul_rf_init"):
+ vax1_0_fused_1 = T.axis.spatial(512, ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads()
T.writes(C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0])
- C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0] =
C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0] + T.Cast("float32", V[0, 0,
vax1_0_fused_0 * 2048 + vax1_0_fused_1 * 8 + vax1_1]) * T.Cast("float32",
(T.Cast("float16", T.bitwise_and(T.shift_right(W[v0, (vax1_0_fused_0 * 2048 +
vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 2048 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * S[v0, (vax1_0_fused_0 * 2048 + vax1_0 [...]
- for ax1_fused in range(1): # pylint: disable=unused-variable
- for ax0_fused_1 in T.thread_binding(256,
thread="threadIdx.x"):
- with T.block("matmul"):
- vax1_0_fused_1, v0 = T.axis.remap("RS",
[ax0_fused_1, ax0_fused])
- T.reads(C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0])
- T.writes(C_fp32_local[0, 0, v0])
- with T.init():
- C_fp32_local[0, 0, v0] = T.float32(0)
- C_fp32_local[0, 0, v0] = C_fp32_local[0, 0, v0] +
C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0]
- with T.block("cast"):
- v0 = T.axis.spatial(4096, ax0_fused)
- T.reads(C_fp32_local[0, 0, v0])
- T.writes(C[0, 0, v0])
- C[0, 0, v0] = T.Cast("float16", C_fp32_local[0, 0, v0])
+ C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0] =
T.float32(0)
+ for ax1_0_fused_0 in range(1):
+ for ax1_1 in range(8):
+ with T.block("matmul_rf_update"):
+ vax1_0_fused_1 = T.axis.spatial(512,
ax1_0_fused_1)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ vax1_0_fused_0 = T.axis.reduce(1,
ax1_0_fused_0)
+ vax1_1 = T.axis.reduce(8, ax1_1)
+ T.reads(C_fp32_rf_local[vax1_0_fused_1, 0,
0, v0], V[0, 0, vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1], W[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 8], S[v0,
(vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1) // 32])
+ T.writes(C_fp32_rf_local[vax1_0_fused_1,
0, 0, v0])
+ C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0]
= C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0] + T.Cast("float32", V[0, 0,
vax1_0_fused_0 * 4096 + vax1_0_fused_1 * 8 + vax1_1]) * T.Cast("float32",
(T.Cast("float16", T.bitwise_and(T.shift_right(W[v0, (vax1_0_fused_0 * 4096 +
vax1_0_fused_1 * 8 + vax1_1) // 8], T.Cast("uint32", (vax1_0_fused_0 * 4096 +
vax1_0_fused_1 * 8 + vax1_1) % 8) * T.uint32(4)), T.uint32(15))) -
T.float16(7)) * S[v0, (vax1_0_fused_0 * 4096 [...]
+ for ax1_fused in range(1):
+ for ax0 in T.thread_binding(512, thread="threadIdx.x"):
+ with T.block("matmul"):
+ vax1_0_fused_1 = T.axis.reduce(512, ax0)
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads(C_fp32_rf_local[vax1_0_fused_1, 0, 0,
v0])
+ T.writes(C_fp32_local[0, 0, v0])
+ with T.init():
+ C_fp32_local[0, 0, v0] = T.float32(0)
+ C_fp32_local[0, 0, v0] = C_fp32_local[0, 0,
v0] + C_fp32_rf_local[vax1_0_fused_1, 0, 0, v0]
+ with T.block("cast"):
+ v0 = T.axis.spatial(4096, ax0_fused)
+ T.reads(C_fp32_local[0, 0, v0])
+ T.writes(C[0, 0, v0])
+ C[0, 0, v0] = T.Cast("float16", C_fp32_local[0, 0, v0])
# fmt: on
@@ -446,44 +501,50 @@ def test_reduction_no_spatial():
@I.ir_module
class After:
@T.prim_func
- def main(A: T.Buffer((1, 1, 4096), "float16"), B: T.Buffer((4096,),
"float16"), rms_norm: T.Buffer((1, 4096), "float16")):
- T.func_attr({"global_symbol": "main", "tir.noalias": True,
"tir.is_scheduled": 1})
- # with T.block("root"):
- Ared_temp_shared = T.alloc_buffer((1, 1), scope="shared")
- Ared_temp_rf_local = T.alloc_buffer((256, 1, 1), scope="local")
- for ax0_fused in T.thread_binding(T.int64(1),
thread="blockIdx.x"): # pylint: disable=unused-variable
- for ax1_fused_1 in T.thread_binding(256, thread="threadIdx.x",
annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}):
- with T.block("Ared_temp_rf_init"):
- vax1_fused_1 = T.axis.spatial(256, ax1_fused_1)
- v0 = T.axis.spatial(T.int64(1), T.int64(0))
- T.reads()
- T.writes(Ared_temp_rf_local[vax1_fused_1, 0, 0])
- Ared_temp_rf_local[vax1_fused_1, 0, 0] = T.float32(0)
- for ax1_fused_0, u in T.grid(16, 1): # pylint:
disable=unused-variable
- with T.block("Ared_temp_rf_update"):
- vax1_fused_1 = T.axis.spatial(256, ax1_fused_1)
+ def main(A_handle: T.handle, B_handle: T.handle, rms_norm_handle:
T.handle):
+ T.func_attr({"tir.is_scheduled": 1, "tir.noalias": T.bool(True)})
+ A = T.match_buffer(A_handle, (1, 1, 4096), "float16")
+ B = T.match_buffer(B_handle, (4096,), "float16")
+ rms_norm = T.match_buffer(rms_norm_handle, (1, 4096), "float16")
+ with T.block("root"):
+ T.reads()
+ T.writes()
+ Ared_temp_shared = T.alloc_buffer((1, 1), scope="shared")
+ Ared_temp_rf_local = T.alloc_buffer((1024, 1, 1),
scope="local")
+ for ax0_fused in T.thread_binding(T.int64(1),
thread="blockIdx.x"):
+ for ax1_fused_1 in T.thread_binding(1024,
thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256,
"pragma_unroll_explicit": 1}):
+ with T.block("Ared_temp_rf_init"):
+ vax1_fused_1 = T.axis.spatial(1024, ax1_fused_1)
v0 = T.axis.spatial(T.int64(1), T.int64(0))
- vax1_fused_0 = T.axis.reduce(16, ax1_fused_0)
- T.reads(Ared_temp_rf_local[vax1_fused_1, 0, 0],
A[0, 0, vax1_fused_0 * 256 + vax1_fused_1])
+ T.reads()
T.writes(Ared_temp_rf_local[vax1_fused_1, 0, 0])
- Ared_temp_rf_local[vax1_fused_1, 0, 0] =
Ared_temp_rf_local[vax1_fused_1, 0, 0] + T.Cast("float32", A[0, 0, vax1_fused_0
* 256 + vax1_fused_1]) * T.Cast("float32", A[0, 0, vax1_fused_0 * 256 +
vax1_fused_1])
- for ax1_fused in range(T.int64(1)): # pylint:
disable=unused-variable
- for ax0 in T.thread_binding(256, thread="threadIdx.x"):
- with T.block("Ared_temp"):
- vax1_fused_1 = T.axis.reduce(256, ax0)
- v0 = T.axis.spatial(T.int64(1), T.int64(0))
- T.reads(Ared_temp_rf_local[vax1_fused_1, 0, 0])
- T.writes(Ared_temp_shared[0, 0])
- with T.init():
- Ared_temp_shared[0, 0] = T.float32(0)
- Ared_temp_shared[0, 0] = Ared_temp_shared[0, 0] +
Ared_temp_rf_local[vax1_fused_1, 0, 0]
- for ax0_fused_0 in range(16):
- for ax0_fused_1 in T.thread_binding(256,
thread="threadIdx.x"):
- with T.block("rms_norm"):
- v0 = T.axis.spatial(4096, ax0_fused_0 * 256 +
ax0_fused_1)
- T.reads(B[v0], A[0, 0, v0], Ared_temp_shared[0, 0])
- T.writes(rms_norm[0, v0])
- rms_norm[0, v0] = T.Cast("float16",
T.Cast("float32", B[v0]) * (T.Cast("float32", A[0, 0, v0]) /
T.sqrt(Ared_temp_shared[0, 0] * T.float32(0.000244140625) +
T.float32(9.9999999999999995e-07))))
+ Ared_temp_rf_local[vax1_fused_1, 0, 0] =
T.float32(0)
+ for ax1_fused_0 in range(4):
+ for u in range(1):
+ with T.block("Ared_temp_rf_update"):
+ vax1_fused_1 = T.axis.spatial(1024,
ax1_fused_1)
+ v0 = T.axis.spatial(T.int64(1), T.int64(0))
+ vax1_fused_0 = T.axis.reduce(4,
ax1_fused_0)
+ T.reads(Ared_temp_rf_local[vax1_fused_1,
0, 0], A[0, 0, vax1_fused_0 * 1024 + vax1_fused_1])
+ T.writes(Ared_temp_rf_local[vax1_fused_1,
0, 0])
+ Ared_temp_rf_local[vax1_fused_1, 0, 0] =
Ared_temp_rf_local[vax1_fused_1, 0, 0] + T.Cast("float32", A[0, 0, vax1_fused_0
* 1024 + vax1_fused_1]) * T.Cast("float32", A[0, 0, vax1_fused_0 * 1024 +
vax1_fused_1])
+ for ax1_fused in range(T.int64(1)):
+ for ax0 in T.thread_binding(1024,
thread="threadIdx.x"):
+ with T.block("Ared_temp"):
+ vax1_fused_1 = T.axis.reduce(1024, ax0)
+ v0 = T.axis.spatial(T.int64(1), T.int64(0))
+ T.reads(Ared_temp_rf_local[vax1_fused_1, 0, 0])
+ T.writes(Ared_temp_shared[0, 0])
+ with T.init():
+ Ared_temp_shared[0, 0] = T.float32(0)
+ Ared_temp_shared[0, 0] = Ared_temp_shared[0,
0] + Ared_temp_rf_local[vax1_fused_1, 0, 0]
+ for ax0_fused_0 in range(4):
+ for ax0_fused_1 in T.thread_binding(1024,
thread="threadIdx.x"):
+ with T.block("rms_norm"):
+ v0 = T.axis.spatial(4096, ax0_fused_0 * 1024 +
ax0_fused_1)
+ T.reads(B[v0], A[0, 0, v0],
Ared_temp_shared[0, 0])
+ T.writes(rms_norm[0, v0])
+ rms_norm[0, v0] = T.Cast("float16",
T.Cast("float32", B[v0]) * (T.Cast("float32", A[0, 0, v0]) /
T.sqrt(Ared_temp_shared[0, 0] * T.float32(0.000244140625) +
T.float32(9.9999999999999995e-07))))
# fmt: on
target = Target("nvidia/geforce-rtx-3090-ti")
with target: