Note
You can click here to run the Jupyter notebook locally.
Transformation#
In this section, we will get to the main ingredients of the compilation flows - transformations of primitive tensor functions.
In the previous section, we have given an example of how to write
mm_relu using TensorIR. In practice, there can be multiple ways to implement
the same functionality, and each implementation can result in different performance.
Note
This tutorial primarily illustrates the application of TensorIR Transformation, rather than delving into optimization techniques.
First, let’s take a look at the implementation of mm_relu in the previous section:
import tvm
from tvm.script import ir as I
from tvm.script import s_tir as Ts
from tvm.script import tirx as T
@I.ir_module
class MyModule:
@Ts.function
def main(
A: T.Tensor((128, 128), "float32"),
B: T.Tensor((128, 128), "float32"),
C: T.Tensor((128, 128), "float32"),
):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
Y = Ts.sblock_alloc_buffer((128, 128))
for i, j, k in T.grid(128, 128, 128):
with Ts.sblock("Y"):
vi, vj, vk = Ts.axis.remap("SSR", [i, j, k])
with Ts.init():
Y[vi, vj] = T.float32(0)
Y[vi, vj] = Y[vi, vj] + A[vi, vk] * B[vk, vj]
for i, j in T.grid(128, 128):
with Ts.sblock("C"):
vi, vj = Ts.axis.remap("SS", [i, j])
C[vi, vj] = T.max(Y[vi, vj], T.float32(0))
Before we transform the function, let’s first evaluate the performance of the original implementation.
import numpy as np
a_np = np.random.uniform(size=(128, 128)).astype("float32")
b_np = np.random.uniform(size=(128, 128)).astype("float32")
c_np = a_np @ b_np
a_nd = tvm.runtime.tensor(a_np)
b_nd = tvm.runtime.tensor(b_np)
c_nd = tvm.runtime.tensor(np.zeros((128, 128), dtype="float32"))
def evaluate(mod: tvm.IRModule):
lib = tvm.tirx.build(mod, target="llvm")
# check correctness
lib(a_nd, b_nd, c_nd)
np.testing.assert_allclose(c_nd.numpy(), c_np, rtol=1e-5)
# evaluate performance
f_timer = lib.time_evaluator("main", tvm.cpu())
print(f_timer(a_nd, b_nd, c_nd))
evaluate(MyModule)
Execution time summary:
mean (ms) median (ms) max (ms) min (ms) std (ms)
2.4568 2.4568 2.4568 2.4568 0.0000
Initialization Schedule#
We initiate the process of code transformation by establishing a Schedule helper class, utilizing the provided MyModule as input.
sch = tvm.s_tir.Schedule(MyModule)
Loop Tiling#
Subsequently, we execute the requisite operations to acquire a reference to block Y and its associated loops.
We now proceed to execute the transformations. The initial modification involves
splitting loop j into two separate loops, with the inner loop possessing a
length of 8. It is crucial to understand that the transformation process is procedural;
thus, inadvertent execution of the block twice will yield an error stating the
non-existence of variable j.
The outcome of the transformation can be examined, as it is retained within sch.mod.
sch.mod.show()
from __future__ import annotations
# from tvm.script import ir as I
# from tvm.script import s_tir as Ts
# from tvm.script import tirx as T
@I.ir_module
class Module:
@Ts.function
def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
v = Ts.sblock_alloc_buffer((128, 128), "float32")
for i in range(128):
for j_0 in range(16):
for j_1 in range(8):
for k in range(128):
with Ts.sblock("Y"):
v_1 = Ts.axis.spatial(128, i)
v_2 = Ts.axis.spatial(128, j_0 * 8 + j_1)
v_3 = Ts.axis.reduce(128, k)
Ts.reads(A[v_1, v_3:v_3 + 1], B[v_3, v_2:v_2 + 1])
Ts.writes(v[v_1, v_2:v_2 + 1])
with Ts.init():
v[v_1, v_2] = T.float32(0.0)
v[v_1, v_2] = v[v_1, v_2] + A[v_1, v_3] * B[v_3, v_2]
for i_1 in range(128):
for j in range(128):
with Ts.sblock("C"):
v_4 = Ts.axis.spatial(128, i_1)
v_5 = Ts.axis.spatial(128, j)
Ts.reads(v[v_4, v_5:v_5 + 1])
Ts.writes(C[v_4, v_5:v_5 + 1])
C[v_4, v_5] = T.max(v[v_4, v_5], T.float32(0.0))
Following the initial transformation phase, two supplementary loops, j_0 and j_1,
have been generated with respective ranges of 16 and 8. The subsequent
action involves reordering these two loops.
from __future__ import annotations
# from tvm.script import ir as I
# from tvm.script import s_tir as Ts
# from tvm.script import tirx as T
@I.ir_module
class Module:
@Ts.function
def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
v = Ts.sblock_alloc_buffer((128, 128), "float32")
for i in range(128):
for j_0 in range(16):
for k in range(128):
for j_1 in range(8):
with Ts.sblock("Y"):
v_1 = Ts.axis.spatial(128, i)
v_2 = Ts.axis.spatial(128, j_0 * 8 + j_1)
v_3 = Ts.axis.reduce(128, k)
Ts.reads(A[v_1, v_3:v_3 + 1], B[v_3, v_2:v_2 + 1])
Ts.writes(v[v_1, v_2:v_2 + 1])
with Ts.init():
v[v_1, v_2] = T.float32(0.0)
v[v_1, v_2] = v[v_1, v_2] + A[v_1, v_3] * B[v_3, v_2]
for i_1 in range(128):
for j in range(128):
with Ts.sblock("C"):
v_4 = Ts.axis.spatial(128, i_1)
v_5 = Ts.axis.spatial(128, j)
Ts.reads(v[v_4, v_5:v_5 + 1])
Ts.writes(C[v_4, v_5:v_5 + 1])
C[v_4, v_5] = T.max(v[v_4, v_5], T.float32(0.0))
Execution time summary:
mean (ms) median (ms) max (ms) min (ms) std (ms)
0.8632 0.8632 0.8632 0.8632 0.0000
Leverage Localities#
Subsequently, we will execute two additional transformation steps to achieve a different variant. First, we employ a primitive known as reverse_compute_at to relocate block C to an inner loop of Y.
from __future__ import annotations
# from tvm.script import ir as I
# from tvm.script import s_tir as Ts
# from tvm.script import tirx as T
@I.ir_module
class Module:
@Ts.function
def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
v = Ts.sblock_alloc_buffer((128, 128), "float32")
for i in range(128):
for j_0 in range(16):
for k in range(128):
for j_1 in range(8):
with Ts.sblock("Y"):
v_1 = Ts.axis.spatial(128, i)
v_2 = Ts.axis.spatial(128, j_0 * 8 + j_1)
v_3 = Ts.axis.reduce(128, k)
Ts.reads(A[v_1, v_3:v_3 + 1], B[v_3, v_2:v_2 + 1])
Ts.writes(v[v_1, v_2:v_2 + 1])
with Ts.init():
v[v_1, v_2] = T.float32(0.0)
v[v_1, v_2] = v[v_1, v_2] + A[v_1, v_3] * B[v_3, v_2]
for ax0 in range(8):
with Ts.sblock("C"):
v_4 = Ts.axis.spatial(128, i)
v_5 = Ts.axis.spatial(128, j_0 * 8 + ax0)
Ts.reads(v[v_4, v_5:v_5 + 1])
Ts.writes(C[v_4, v_5:v_5 + 1])
C[v_4, v_5] = T.max(v[v_4, v_5], T.float32(0.0))
Rewrite Reduction#
Until now, the reduction initialization and update step have been maintained together
within a single block body. This amalgamated form facilitates loop transformations,
as the outer loops i, j of initialization and updates generally need to remain
synchronized.
Following the loop transformations, we can segregate the initialization of Y’s elements from the reduction update via the decompose_reduction primitive.
from __future__ import annotations
# from tvm.script import ir as I
# from tvm.script import s_tir as Ts
# from tvm.script import tirx as T
@I.ir_module
class Module:
@Ts.function
def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
v = Ts.sblock_alloc_buffer((128, 128), "float32")
for i in range(128):
for j_0 in range(16):
for j_1_init in range(8):
with Ts.sblock("Y_init"):
v_1 = Ts.axis.spatial(128, i)
v_2 = Ts.axis.spatial(128, j_0 * 8 + j_1_init)
Ts.reads()
Ts.writes(v[v_1, v_2:v_2 + 1])
v[v_1, v_2] = T.float32(0.0)
for k in range(128):
for j_1 in range(8):
with Ts.sblock("Y_update"):
v_3 = Ts.axis.spatial(128, i)
v_4 = Ts.axis.spatial(128, j_0 * 8 + j_1)
v_5 = Ts.axis.reduce(128, k)
Ts.reads(v[v_3, v_4:v_4 + 1], A[v_3, v_5:v_5 + 1], B[v_5, v_4:v_4 + 1])
Ts.writes(v[v_3, v_4:v_4 + 1])
v[v_3, v_4] = v[v_3, v_4] + A[v_3, v_5] * B[v_5, v_4]
for ax0 in range(8):
with Ts.sblock("C"):
v_6 = Ts.axis.spatial(128, i)
v_7 = Ts.axis.spatial(128, j_0 * 8 + ax0)
Ts.reads(v[v_6, v_7:v_7 + 1])
Ts.writes(C[v_6, v_7:v_7 + 1])
C[v_6, v_7] = T.max(v[v_6, v_7], T.float32(0.0))
Execution time summary:
mean (ms) median (ms) max (ms) min (ms) std (ms)
0.3482 0.3482 0.3482 0.3482 0.0000
Trace the Transformation#
TensorIR schedule is a procedural language, and the transformation is executed in a step-by-step manner. We can trace the transformation by printing the schedule or the history of the schedule.
We’ve already see the schedule by printing sch.mod. We can also print the history
of the schedule by sch.trace.
sch.trace.show()
# from tvm import s_tir
def apply_trace(sch: s_tir.Schedule) -> None:
b0 = sch.get_sblock(name="Y", func_name="main")
l1, l2, l3 = sch.get_loops(block=b0)
l4, l5 = sch.split(loop=l2, factors=[None, 8], preserve_unit_iters=True, disable_predication=False)
sch.reorder(l4, l3, l5)
b6 = sch.get_sblock(name="C", func_name="main")
sch.reverse_compute_at(block=b6, loop=l4, preserve_unit_loops=False, index=-1)
b7 = sch.decompose_reduction(block=b0, loop=l3)
Alternatively, we can output the IRModule in conjunction with the historical trace.
sch.show()
from __future__ import annotations
# from tvm.script import ir as I
# from tvm.script import s_tir as Ts
# from tvm.script import tirx as T
@I.ir_module
class Module:
@Ts.function
def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")):
T.func_attr({"tirx.noalias": True})
with Ts.sblock("root"):
Ts.reads()
Ts.writes()
v = Ts.sblock_alloc_buffer((128, 128), "float32")
for i in range(128):
for j_0 in range(16):
for j_1_init in range(8):
with Ts.sblock("Y_init"):
v_1 = Ts.axis.spatial(128, i)
v_2 = Ts.axis.spatial(128, j_0 * 8 + j_1_init)
Ts.reads()
Ts.writes(v[v_1, v_2:v_2 + 1])
v[v_1, v_2] = T.float32(0.0)
for k in range(128):
for j_1 in range(8):
with Ts.sblock("Y_update"):
v_3 = Ts.axis.spatial(128, i)
v_4 = Ts.axis.spatial(128, j_0 * 8 + j_1)
v_5 = Ts.axis.reduce(128, k)
Ts.reads(v[v_3, v_4:v_4 + 1], A[v_3, v_5:v_5 + 1], B[v_5, v_4:v_4 + 1])
Ts.writes(v[v_3, v_4:v_4 + 1])
v[v_3, v_4] = v[v_3, v_4] + A[v_3, v_5] * B[v_5, v_4]
for ax0 in range(8):
with Ts.sblock("C"):
v_6 = Ts.axis.spatial(128, i)
v_7 = Ts.axis.spatial(128, j_0 * 8 + ax0)
Ts.reads(v[v_6, v_7:v_7 + 1])
Ts.writes(C[v_6, v_7:v_7 + 1])
C[v_6, v_7] = T.max(v[v_6, v_7], T.float32(0.0))
# from tvm import s_tir
def apply_trace(sch: s_tir.Schedule) -> None:
b0 = sch.get_sblock(name="Y", func_name="main")
l1, l2, l3 = sch.get_loops(block=b0)
l4, l5 = sch.split(loop=l2, factors=[None, 8], preserve_unit_iters=True, disable_predication=False)
sch.reorder(l4, l3, l5)
b6 = sch.get_sblock(name="C", func_name="main")
sch.reverse_compute_at(block=b6, loop=l4, preserve_unit_loops=False, index=-1)
b7 = sch.decompose_reduction(block=b0, loop=l3)