18 KiB
18 KiB
In [1]:
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import from_dlpack
import numpy as np
import torchIn [2]:
@cute.jit
def load_and_store(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):
"""
Load data from memory and store the result to memory.
:param res: The destination tensor to store the result.
:param a: The source tensor to be loaded.
:param b: The source tensor to be loaded.
"""
a_vec = a.load()
print(f"a_vec: {a_vec}") # prints `a_vec: vector<12xf32> o (3, 4)`
b_vec = b.load()
print(f"b_vec: {b_vec}") # prints `b_vec: vector<12xf32> o (3, 4)`
res.store(a_vec + b_vec)
cute.print_tensor(res)
a = np.ones(12).reshape((3, 4)).astype(np.float32)
b = np.ones(12).reshape((3, 4)).astype(np.float32)
c = np.zeros(12).reshape((3, 4)).astype(np.float32)
load_and_store(from_dlpack(c), from_dlpack(a), from_dlpack(b))a_vec: tensor_value<vector<12xf32> o (3, 4)>
b_vec: tensor_value<vector<12xf32> o (3, 4)>
tensor(raw_ptr(0x0000000006cff170: f32, generic, align<4>) o (3,4):(4,1), data=
[[ 2.000000, 2.000000, 2.000000, 2.000000, ],
[ 2.000000, 2.000000, 2.000000, 2.000000, ],
[ 2.000000, 2.000000, 2.000000, 2.000000, ]])
In [3]:
@cute.jit
def apply_slice(src: cute.Tensor, dst: cute.Tensor, indices: cutlass.Constexpr):
"""
Apply slice operation on the src tensor and store the result to the dst tensor.
:param src: The source tensor to be sliced.
:param dst: The destination tensor to store the result.
:param indices: The indices to slice the source tensor.
"""
src_vec = src.load()
dst_vec = src_vec[indices]
print(f"{src_vec} -> {dst_vec}")
if cutlass.const_expr(isinstance(dst_vec, cute.TensorSSA)):
dst.store(dst_vec)
cute.print_tensor(dst)
else:
dst[0] = dst_vec
cute.print_tensor(dst)
def slice_1():
src_shape = (4, 2, 3)
dst_shape = (4, 3)
indices = (None, 1, None)
"""
a:
[[[ 0. 1. 2.]
[ 3. 4. 5.]]
[[ 6. 7. 8.]
[ 9. 10. 11.]]
[[12. 13. 14.]
[15. 16. 17.]]
[[18. 19. 20.]
[21. 22. 23.]]]
"""
a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)
dst = np.random.randn(*dst_shape).astype(np.float32)
apply_slice(from_dlpack(a), from_dlpack(dst), indices)
slice_1()tensor_value<vector<24xf32> o (4, 2, 3)> -> tensor_value<vector<12xf32> o (4, 3)>
tensor(raw_ptr(0x00000000071acaf0: f32, generic, align<4>) o (4,3):(3,1), data=
[[ 3.000000, 4.000000, 5.000000, ],
[ 9.000000, 10.000000, 11.000000, ],
[ 15.000000, 16.000000, 17.000000, ],
[ 21.000000, 22.000000, 23.000000, ]])
In [4]:
def slice_2():
src_shape = (4, 2, 3)
dst_shape = (1,)
indices = 10
a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)
dst = np.random.randn(*dst_shape).astype(np.float32)
apply_slice(from_dlpack(a), from_dlpack(dst), indices)
slice_2()tensor_value<vector<24xf32> o (4, 2, 3)> -> ?
tensor(raw_ptr(0x00000000013cbbe0: f32, generic, align<4>) o (1):(1), data=
[ 10.000000, ])
In [5]:
@cute.jit
def binary_op_1(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):
a_vec = a.load()
b_vec = b.load()
add_res = a_vec + b_vec
res.store(add_res)
cute.print_tensor(res) # prints [3.000000, 3.000000, 3.000000]
sub_res = a_vec - b_vec
res.store(sub_res)
cute.print_tensor(res) # prints [-1.000000, -1.000000, -1.000000]
mul_res = a_vec * b_vec
res.store(mul_res)
cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]
div_res = a_vec / b_vec
res.store(div_res)
cute.print_tensor(res) # prints [0.500000, 0.500000, 0.500000]
floor_div_res = a_vec // b_vec
res.store(floor_div_res)
cute.print_tensor(res) # prints [0.000000, 0.000000, 0.000000]
mod_res = a_vec % b_vec
res.store(mod_res)
cute.print_tensor(res) # prints [1.000000, 1.000000, 1.000000]
a = np.empty((3,), dtype=np.float32)
a.fill(1.0)
b = np.empty((3,), dtype=np.float32)
b.fill(2.0)
res = np.empty((3,), dtype=np.float32)
binary_op_1(from_dlpack(res), from_dlpack(a), from_dlpack(b))tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[ 3.000000, ],
[ 3.000000, ],
[ 3.000000, ])
tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[-1.000000, ],
[-1.000000, ],
[-1.000000, ])
tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[ 2.000000, ],
[ 2.000000, ],
[ 2.000000, ])
tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[ 0.500000, ],
[ 0.500000, ],
[ 0.500000, ])
tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[ 0.000000, ],
[ 0.000000, ],
[ 0.000000, ])
tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=
[ 1.000000, ],
[ 1.000000, ],
[ 1.000000, ])
In [6]:
@cute.jit
def binary_op_2(res: cute.Tensor, a: cute.Tensor, c: cutlass.Constexpr):
a_vec = a.load()
add_res = a_vec + c
res.store(add_res)
cute.print_tensor(res) # prints [3.000000, 3.000000, 3.000000]
sub_res = a_vec - c
res.store(sub_res)
cute.print_tensor(res) # prints [-1.000000, -1.000000, -1.000000]
mul_res = a_vec * c
res.store(mul_res)
cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]
div_res = a_vec / c
res.store(div_res)
cute.print_tensor(res) # prints [0.500000, 0.500000, 0.500000]
floor_div_res = a_vec // c
res.store(floor_div_res)
cute.print_tensor(res) # prints [0.000000, 0.000000, 0.000000]
mod_res = a_vec % c
res.store(mod_res)
cute.print_tensor(res) # prints [1.000000, 1.000000, 1.000000]
a = np.empty((3,), dtype=np.float32)
a.fill(1.0)
c = 2.0
res = np.empty((3,), dtype=np.float32)
binary_op_2(from_dlpack(res), from_dlpack(a), c)tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[ 3.000000, ],
[ 3.000000, ],
[ 3.000000, ])
tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[-1.000000, ],
[-1.000000, ],
[-1.000000, ])
tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[ 2.000000, ],
[ 2.000000, ],
[ 2.000000, ])
tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[ 0.500000, ],
[ 0.500000, ],
[ 0.500000, ])
tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[ 0.000000, ],
[ 0.000000, ],
[ 0.000000, ])
tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=
[ 1.000000, ],
[ 1.000000, ],
[ 1.000000, ])
In [7]:
@cute.jit
def binary_op_3(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):
a_vec = a.load()
b_vec = b.load()
gt_res = a_vec > b_vec
res.store(gt_res)
"""
ge_res = a_ >= b_ # [False, True, False]
lt_res = a_ < b_ # [True, False, True]
le_res = a_ <= b_ # [True, False, True]
eq_res = a_ == b_ # [False, False, False]
"""
a = np.array([1, 2, 3], dtype=np.float32)
b = np.array([2, 1, 4], dtype=np.float32)
res = np.empty((3,), dtype=np.bool_)
binary_op_3(from_dlpack(res), from_dlpack(a), from_dlpack(b))
print(res) # prints [False, True, False]
[False True False]
In [8]:
@cute.jit
def binary_op_4(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):
a_vec = a.load()
b_vec = b.load()
xor_res = a_vec ^ b_vec
res.store(xor_res)
# or_res = a_vec | b_vec
# res.store(or_res) # prints [3, 2, 7]
# and_res = a_vec & b_vec
# res.store(and_res) # prints [0, 2, 0]
a = np.array([1, 2, 3], dtype=np.int32)
b = np.array([2, 2, 4], dtype=np.int32)
res = np.empty((3,), dtype=np.int32)
binary_op_4(from_dlpack(res), from_dlpack(a), from_dlpack(b))
print(res) # prints [3, 0, 7][3 0 7]
In [9]:
@cute.jit
def unary_op_1(res: cute.Tensor, a: cute.Tensor):
a_vec = a.load()
sqrt_res = cute.math.sqrt(a_vec)
res.store(sqrt_res)
cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]
sin_res = cute.math.sin(a_vec)
res.store(sin_res)
cute.print_tensor(res) # prints [-0.756802, -0.756802, -0.756802]
exp2_res = cute.math.exp2(a_vec)
res.store(exp2_res)
cute.print_tensor(res) # prints [16.000000, 16.000000, 16.000000]
a = np.array([4.0, 4.0, 4.0], dtype=np.float32)
res = np.empty((3,), dtype=np.float32)
unary_op_1(from_dlpack(res), from_dlpack(a))tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=
[ 2.000000, ],
[ 2.000000, ],
[ 2.000000, ])
tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=
[-0.756802, ],
[-0.756802, ],
[-0.756802, ])
tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=
[ 16.000000, ],
[ 16.000000, ],
[ 16.000000, ])
In [10]:
@cute.jit
def reduction_op(a: cute.Tensor):
"""
Apply reduction operation on the src tensor.
:param src: The source tensor to be reduced.
"""
a_vec = a.load()
red_res = a_vec.reduce(
cute.ReductionOp.ADD,
0.0,
reduction_profile=0
)
cute.printf(red_res) # prints 21.000000
red_res = a_vec.reduce(
cute.ReductionOp.ADD,
0.0,
reduction_profile=(None, 1)
)
# We can't print the TensorSSA directly at this point, so we store it to a new Tensor and print it.
res = cute.make_fragment(red_res.shape, cutlass.Float32)
res.store(red_res)
cute.print_tensor(res) # prints [6.000000, 15.000000]
red_res = a_vec.reduce(
cute.ReductionOp.ADD,
1.0,
reduction_profile=(1, None)
)
res = cute.make_fragment(red_res.shape, cutlass.Float32)
res.store(red_res)
cute.print_tensor(res) # prints [6.000000, 8.000000, 10.000000]
a = np.array([[1, 2, 3], [4, 5, 6]], dtype=np.float32)
reduction_op(from_dlpack(a))21.000000
tensor(raw_ptr(0x00007ffd1ea2bca0: f32, rmem, align<32>) o (2):(1), data=
[ 6.000000, ],
[ 15.000000, ])
tensor(raw_ptr(0x00007ffd1ea2bcc0: f32, rmem, align<32>) o (3):(1), data=
[ 6.000000, ],
[ 8.000000, ],
[ 10.000000, ])