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"):

Reply via email to