mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-03 23:08:48 +08:00
49 lines
1.5 KiB
Python
49 lines
1.5 KiB
Python
import torch
|
|
from tvm import tl
|
|
import tvm.tl.language as T
|
|
|
|
|
|
def transpose(M, N):
|
|
dtype = "float"
|
|
BLK = 64
|
|
|
|
@T.prim_func
|
|
def main(ins: T.Buffer((M, N), dtype), outs: T.Buffer((N, M), dtype)):
|
|
with T.Kernel(T.ceildiv(M, BLK), T.ceildiv(N, BLK), threads=256) as (bx, by):
|
|
shared = T.alloc_shared((BLK, BLK), dtype)
|
|
local = T.alloc_fragment((BLK, BLK), dtype)
|
|
local_t = T.alloc_fragment((BLK, BLK), dtype)
|
|
T.annotate_layout(
|
|
{
|
|
# pad by 4
|
|
shared: T.Layout(shared.shape, lambda i, j: i * (BLK + 4) + j),
|
|
# assign 4x4 float tile to each thread
|
|
local: T.Fragment(local.shape, lambda i, j: j // 4 + 16 * (i // 4)),
|
|
}
|
|
)
|
|
T.copy(ins[BLK * by, BLK * bx], shared)
|
|
T.copy(shared, local)
|
|
for i, j in T.Parallel(BLK, BLK):
|
|
local_t[i, j] = local[j, i]
|
|
T.copy(local_t, shared)
|
|
T.copy(shared, outs[BLK * bx, BLK * by])
|
|
|
|
return main
|
|
|
|
|
|
def ref_program(A):
|
|
return A.T.contiguous()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
M, N = 8192, 8192
|
|
program = transpose(M, N)
|
|
mod, params = tl.lower(program)
|
|
mod = tl.Profiler(mod, params, [1], tl.TensorSupplyType.Integer)
|
|
mod.assert_allclose(ref_program)
|
|
|
|
latency = mod.do_bench(ref_program, warmup=500)
|
|
print("{:.2f} ms".format(latency))
|
|
latency = mod.do_bench(mod.func)
|
|
print("{:.2f} ms".format(latency))
|