tvm import hook

This commit is contained in:
2023-09-28 09:24:32 -07:00
parent adab724caa
commit c36d0e3bd8
+29 -25
View File
@@ -1,26 +1,30 @@
# https://tvm.apache.org/docs/tutorial/tensor_expr_get_started.html#example-2-manually-optimizing-matrix-multiplication-with-te
import tvm
from tvm import te
#print(tvm.target.Target.list_kinds())
M, N, K = 1024, 1024, 1024
# c, opencl
target = tvm.target.Target(target="c")
try:
import tvm
from tvm import te
#print(tvm.target.Target.list_kinds())
# TVM Matrix Multiplication using TE
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")
# c, opencl
target = tvm.target.Target(target="c")
# Default schedule
s = te.create_schedule(C.op)
#print(tvm.lower(s, [A, B, C], simple_mode=True))
# TVM Matrix Multiplication using TE
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")
# Output C code
func = tvm.build(s, [A, B, C], target=target, name="mmult")
print(func.get_source())
# Default schedule
s = te.create_schedule(C.op)
#print(tvm.lower(s, [A, B, C], simple_mode=True))
# Output C code
func = tvm.build(s, [A, B, C], target=target, name="mmult")
print(func.get_source())
except ImportError:
print("** please install TVM for TVM output")
# tinygrad version
@@ -35,12 +39,12 @@ A = Tensor.rand(M, K, device="clang")
B = Tensor.rand(K, N, device="clang")
C = (A.reshape(M, 1, K) * B.permute(1,0).reshape(1, N, K)).sum(axis=2)
# capture the kernel. TODO: https://github.com/tinygrad/tinygrad/issues/1812
from tinygrad.jit import CacheCollector
CacheCollector.start()
C.realize()
result = CacheCollector.finish()
print(result[0][0].prg)
sched = C.lazydata.schedule()
from tinygrad.codegen.linearizer import Linearizer
from tinygrad.codegen.kernel import LinearizerOptions
lin = Linearizer(sched[-1][0], LinearizerOptions(has_local=False, supports_float4=False))
#lin.hand_coded_optimizations()
lin.linearize()
from tinygrad.runtime.ops_clang import renderer
src = renderer("mmult", lin.uops)
print(src)