File size: 591 Bytes
8c8128e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 | # Copyright (c) 2025, Wentao Guo, Ted Zadouri, Tri Dao.
from typing import Optional
import cutlass.cute as cute
def make_fake_tensor(dtype, shape, divisibility=1, leading_dim=-1) -> Optional[cute.Tensor]:
if leading_dim < 0:
leading_dim = len(shape) + leading_dim
if dtype is None:
return None
stride = tuple(
cute.sym_int64(divisibility=divisibility) if i != leading_dim else 1
for i in range(len(shape))
)
return cute.runtime.make_fake_tensor(
dtype, shape, stride=stride, assumed_align=divisibility * dtype.width // 8
)
|