In [2]:
import tvm
import tvm.testing
from tvm import te
import numpy
import timeit

# 矩阵的大小
# (M, K) x (K, N)
# 可尝试不同的 shape，TVM 优化的性能有时比 numpy + MKL 更好
M = 1024
K = 1024
N = 1024

# TVM 默认张量数据类型
dtype = "float32"

# 你可能想调整 target 使其和你的任何 CPU 向量扩展匹配
# 例如，如果你为 SIMD 用的是 Intel AVX2（高级向量扩展）ISA，把下面这行换成 `llvm -mcpu=core-avx2` 可以取得最佳性能（或者你所用 CPU 的具体类型）
# 记住你用的是 llvm, 可以用 `llc --version` 命令来获取 CPU 类型，也可以查看 `/proc/cpuinfo` 来获取你处理器支持的更多扩展

target = tvm.target.Target(target="llvm", host="llvm")
dev = tvm.device(target.kind.name, 0)

# 为测试随机生成的张量
a = tvm.nd.array(numpy.random.rand(M, K).astype(dtype), dev)
b = tvm.nd.array(numpy.random.rand(K, N).astype(dtype), dev)

# 重复执行矩阵乘法以获得默认 numpy 实现的性能基线
np_repeat = 100
np_running_time = timeit.timeit(
    setup="import numpy\n"
    "M = " + str(M) + "\n"
    "K = " + str(K) + "\n"
    "N = " + str(N) + "\n"
    'dtype = "float32"\n'
    "a = numpy.random.rand(M, K).astype(dtype)\n"
    "b = numpy.random.rand(K, N).astype(dtype)\n",
    stmt="answer = numpy.dot(a, b)",
    number=np_repeat,
)
print("Numpy running time: %f" % (np_running_time / np_repeat))

answer = numpy.dot(a.numpy(), b.numpy())

Numpy running time: 0.009595


In [3]:
# 用 TE 的 TVM 矩阵乘法
k = te.reduce_axis((0, K), "k")
A = te.placeholder((M, K), name="A")
B = te.placeholder((K, N), name="B")
C = te.compute((M, N), lambda x, y: te.sum(A[x, k] * B[k, y], axis=k), name="C")

# 默认 schedule
s = te.create_schedule(C.op)
func = tvm.build(s, [A, B, C], target=target, name="mmult")

c = tvm.nd.array(numpy.zeros((M, N), dtype=dtype), dev)
func(a, b, c)
tvm.testing.assert_allclose(c.numpy(), answer, rtol=1e-5)

def evaluate_operation(s, vars, target, name, optimization, log):
    func = tvm.build(s, [A, B, C], target=target, name="mmult")
    assert func

    c = tvm.nd.array(numpy.zeros((M, N), dtype=dtype), dev)
    func(a, b, c)
    tvm.testing.assert_allclose(c.numpy(), answer, rtol=1e-5)

    evaluator = func.time_evaluator(func.entry_name, dev, number=10)
    mean_time = evaluator(a, b, c).mean
    print("%s: %f" % (optimization, mean_time))
    log.append((optimization, mean_time))

log = []

evaluate_operation(s, [A, B, C], target=target, name="mmult", optimization="none", log=log)

[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:51:21] /home/zhmi/project/tvm/src/ir/transform.cc:440

[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:51:23] /home/zhmi/project/tvm/src/ir/transform.cc:440

none: 1.738002


In [5]:
print(tvm.lower(s, [A, B, C], simple_mode=True))

# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        for x, y in T.grid(1024, 1024):
            C_1 = T.Buffer((1048576,), data=C.data)
            C_1[x * 1024 + y] = T.float32(0)
            for k in range(1024):
                cse_var_2: T.int32 = x * 1024
                cse_var_1: T.int32 = cse_var_2 + y
                A_1 = T.Buffer((1048576,), data=A.data)
                B_1 = T.Buffer((1048576,), data=B.data)
                C_1[cse_var_1] = C_1[cse_var_1] + A_1[cse_var_2 + k] * B_1[k * 1024 + y]


[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:52:17] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [None]:
#优化一：块优化

In [6]:
bn = 32

# 通过循环切分实现块级化
xo, yo, xi, yi = s[C].tile(C.op.axis[0], C.op.axis[1], bn, bn)
(k,) = s[C].op.reduce_axis
ko, ki = s[C].split(k, factor=4)

# 将归约域提升到块循环外
s[C].reorder(xo, yo, ko, ki, xi, yi)

evaluate_operation(s, [A, B, C], target=target, name="mmult", optimization="blocking", log=log)

[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:00] /home/zhmi/project/tvm/src/ir/transform.cc:440

blocking: 0.271633


In [8]:
print(tvm.lower(s, [A, B, C], simple_mode=True))

# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        for x_outer, y_outer in T.grid(32, 32):
            C_1 = T.Buffer((1048576,), data=C.data)
            for x_inner_init, y_inner_init in T.grid(32, 32):
                C_1[x_outer * 32768 + x_inner_init * 1024 + y_outer * 32 + y_inner_init] = T.float32(0)
            for k_outer, k_inner, x_inner, y_inner in T.grid(256, 4, 32, 32):
                cse_var_3: T.int32 = y_outer * 32
                cse_var_2: T.int32 = x_outer * 32768 + x_inner * 1024
                cse_var_1: T.int32 = cse_var_2 + cse_var_3 + y_inner
                A_1 = T.Buffer((1048576,), data=A.data)
                B_1 = T.Buffer((1048576,), data=B.data

[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:12] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [None]:
#优化二：向量化

In [9]:
# Apply the vectorization optimization
s[C].vectorize(yi)

evaluate_operation(s, [A, B, C], target=target, name="mmult", optimization="vectorization", log=log)

# The generalized IR after vectorization
print(tvm.lower(s, [A, B, C], simple_mode=True))

[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:31] /home/zhmi/project/tvm/src/ir/transform.cc:440

vectorization: 0.277981
# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        for x_outer, y_outer in T.grid(32, 32):
            C_1 = T.Buffer((1048576,), data=C.data)
            for x_inner_init in range(32):
                C_1[x_outer * 32768 + x_inner_init * 1024 + y_outer * 32:x_outer * 32768 + x_inner_init * 1024 + y_outer * 32 + 32] = T.Broadcast(T.float32(0), 32)
            for k_outer, k_inner, x_inner in T.grid(256, 4, 32):
                cse_var_3: T.int32 = y_outer * 32
                cse_var_2: T.int32 = x_outer * 32768 + x_inner * 1024
                cse_var_1: T.int32 = cse_var_2 + cse_var_3
                A_1 = T.Buffer((1048576,), data=A.data)
            

[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:35] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [None]:
#优化三：循环置换

In [10]:
s = te.create_schedule(C.op)
xo, yo, xi, yi = s[C].tile(C.op.axis[0], C.op.axis[1], bn, bn)
(k,) = s[C].op.reduce_axis
ko, ki = s[C].split(k, factor=4)

# re-ordering
# 重新排序
s[C].reorder(xo, yo, ko, xi, ki, yi)
s[C].vectorize(yi)

evaluate_operation(
    s, [A, B, C], target=target, name="mmult", optimization="loop permutation", log=log
)

# Again, print the new generalized IR
# 再次打印新生成的 IR
print(tvm.lower(s, [A, B, C], simple_mode=True))

[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:51] /home/zhmi/project/tvm/src/ir/transform.cc:440

loop permutation: 0.096827
# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        for x_outer, y_outer in T.grid(32, 32):
            C_1 = T.Buffer((1048576,), data=C.data)
            for x_inner_init in range(32):
                C_1[x_outer * 32768 + x_inner_init * 1024 + y_outer * 32:x_outer * 32768 + x_inner_init * 1024 + y_outer * 32 + 32] = T.Broadcast(T.float32(0), 32)
            for k_outer, x_inner, k_inner in T.grid(256, 32, 4):
                cse_var_3: T.int32 = y_outer * 32
                cse_var_2: T.int32 = x_outer * 32768 + x_inner * 1024
                cse_var_1: T.int32 = cse_var_2 + cse_var_3
                A_1 = T.Buffer((1048576,), data=A.data)
         

[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:53:52] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [11]:
#优化四：数组打包
# We have to re-write the algorithm slightly.
# 我们必须稍作改动以重写算法。
packedB = te.compute((N / bn, K, bn), lambda x, y, z: B[y, x * bn + z], name="packedB")
C = te.compute(
    (M, N),
    lambda x, y: te.sum(A[x, k] * packedB[y // bn, k, tvm.tir.indexmod(y, bn)], axis=k),
    name="C",
)

s = te.create_schedule(C.op)

xo, yo, xi, yi = s[C].tile(C.op.axis[0], C.op.axis[1], bn, bn)
(k,) = s[C].op.reduce_axis
ko, ki = s[C].split(k, factor=4)

s[C].reorder(xo, yo, ko, xi, ki, yi)
s[C].vectorize(yi)

x, y, z = s[packedB].op.axis
s[packedB].vectorize(z)
s[packedB].parallel(x)

evaluate_operation(s, [A, B, C], target=target, name="mmult", optimization="array packing", log=log)

# Here is the generated IR after array packing.
# 数组打包后生成的 IR。
print(tvm.lower(s, [A, B, C], simple_mode=True))

[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:11] /home/zhmi/project/tvm/src/ir/transform.cc:440

array packing: 0.114685
# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        packedB = T.allocate([32768], "float32x32", "global")
        packedB_1 = T.Buffer((32768,), "float32x32", data=packedB)
        for x in T.parallel(32):
            for y in range(1024):
                B_1 = T.Buffer((1048576,), data=B.data)
                packedB_1[x * 1024 + y] = B_1[y * 1024 + x * 32:y * 1024 + x * 32 + 32]
        for x_outer, y_outer in T.grid(32, 32):
            C_1 = T.Buffer((1048576,), data=C.data)
            for x_inner_init in range(32):
                C_1[x_outer * 32768 + x_inner_init * 1024 + y_outer * 32:x_outer * 32768 + x_inner_init * 1024 + y_outer * 32 + 32] = T.

[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:12] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [12]:
#优化五：通过缓存优化块写入

In [13]:
s = te.create_schedule(C.op)

# Allocate write cache
# 分配写缓存
CC = s.cache_write(C, "global")

xo, yo, xi, yi = s[C].tile(C.op.axis[0], C.op.axis[1], bn, bn)

# Write cache is computed at yo
# 写缓存在 yo 处被计算
s[CC].compute_at(s[C], yo)

# New inner axes
# 新的内部轴
xc, yc = s[CC].op.axis

(k,) = s[CC].op.reduce_axis
ko, ki = s[CC].split(k, factor=4)
s[CC].reorder(ko, xc, ki, yc)
s[CC].unroll(ki)
s[CC].vectorize(yc)

x, y, z = s[packedB].op.axis
s[packedB].vectorize(z)
s[packedB].parallel(x)

evaluate_operation(s, [A, B, C], target=target, name="mmult", optimization="block caching", log=log)

# Here is the generated IR after write cache blocking.
# 写缓存块级化后生成的 IR。
print(tvm.lower(s, [A, B, C], simple_mode=True))

[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:36] /home/zhmi/project/tvm/src/ir/transform.cc:440

block caching: 0.113304
# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        packedB = T.allocate([32768], "float32x32", "global")
        C_global = T.allocate([1024], "float32", "global")
        packedB_1 = T.Buffer((32768,), "float32x32", data=packedB)
        for x in T.parallel(32):
            for y in range(1024):
                B_1 = T.Buffer((1048576,), data=B.data)
                packedB_1[x * 1024 + y] = B_1[y * 1024 + x * 32:y * 1024 + x * 32 + 32]
        for x_outer, y_outer in T.grid(32, 32):
            C_global_1 = T.Buffer((1024,), data=C_global)
            for x_c_init in range(32):
                C_global_1[x_c_init * 32:x_c_init * 32 + 32] = T.Broadcast(

[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:37] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [14]:
#优化六：并行化

In [15]:
# parallel
# 并行化
s[C].parallel(xo)

x, y, z = s[packedB].op.axis
s[packedB].vectorize(z)
s[packedB].parallel(x)

evaluate_operation(
    s, [A, B, C], target=target, name="mmult", optimization="parallelization", log=log
)

# Here is the generated IR after parallelization.
# 并行化后生成的 IR。
print(tvm.lower(s, [A, B, C], simple_mode=True))

[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:53] /home/zhmi/project/tvm/src/ir/transform.cc:440

parallelization: 0.037634
# from tvm.script import ir as I
# from tvm.script import tir as T

@I.ir_module
class Module:
    @T.prim_func
    def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")):
        T.func_attr({"from_legacy_te_schedule": T.bool(True), "global_symbol": "main", "tir.noalias": T.bool(True)})
        packedB = T.allocate([32768], "float32x32", "global")
        packedB_1 = T.Buffer((32768,), "float32x32", data=packedB)
        for x in T.parallel(32):
            for y in range(1024):
                B_1 = T.Buffer((1048576,), data=B.data)
                packedB_1[x * 1024 + y] = B_1[y * 1024 + x * 32:y * 1024 + x * 32 + 32]
        for x_outer in T.parallel(32):
            C_global = T.allocate([1024], "float32", "global")
            for y_outer in range(32):
                C_global_1 = T.Buffer((1024,), data=C_global)
                for x_c_init in range(32):
                    C_global_1[x

[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.InjectPrefetch
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.TextureFlatten
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlatten
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferShapeLegalize
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferStrideLegalize
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ThreadScopePropagate
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.BufferBindUnwrapper
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.ApplyLayoutTransforms
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.StorageFlattener
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440: Running pass tir.AssertSimplifier
[19:54:54] /home/zhmi/project/tvm/src/ir/transform.cc:440

In [None]:
#矩阵乘法示例总结

In [16]:
baseline = log[0][1]
print("%s\t%s\t%s" % ("Operator".rjust(20), "Timing".rjust(20), "Performance".rjust(20)))
for result in log:
    print(
        "%s\t%s\t%s"
        % (result[0].rjust(20), str(result[1]).rjust(20), str(result[1] / baseline).rjust(20))
    )

            Operator	              Timing	         Performance
                none	  1.7380016103000002	                 1.0
            blocking	        0.2716328607	 0.15629033891005017
       vectorization	 0.27798108229999996	  0.1599429371368747
    loop permutation	        0.0968271348	 0.05571176357154609
       array packing	 0.11468537510000001	 0.06598692108242864
       block caching	        0.1133037901	 0.06519199374069762
     parallelization	           0.0376339	0.021653547256209923
