Fix examples and pytest, run ruff (#3230)
Mark inactive issues and pull requests / mark-inactive-30d (push) Has been cancelled
Mark inactive issues and pull requests / mark-inactive-90d (push) Has been cancelled

Co-authored-by: dePaul Miller <23461061+depaulmillz@users.noreply.github.com>
This commit is contained in:
dePaul Miller
2026-05-21 11:05:38 +08:00
committed by GitHub
co-authored by dePaul Miller
parent 982cb9e718
commit 546c3efa89
11 changed files with 342 additions and 159 deletions
+14 -6
View File
@@ -36,8 +36,10 @@ from cutlass.cute.runtime import from_dlpack
@cute.kernel
def _unary_ops_kernel(
absf_inp: cute.Tensor, absf_out: cute.Tensor,
floor_inp: cute.Tensor, floor_out: cute.Tensor,
absf_inp: cute.Tensor,
absf_out: cute.Tensor,
floor_inp: cute.Tensor,
floor_out: cute.Tensor,
):
tidx, _, _ = cute.arch.thread_idx()
absf_out[tidx] = cute.math.absf(absf_inp[tidx])
@@ -46,8 +48,10 @@ def _unary_ops_kernel(
@cute.jit
def _unary_ops_host(
absf_inp: cute.Tensor, absf_out: cute.Tensor,
floor_inp: cute.Tensor, floor_out: cute.Tensor,
absf_inp: cute.Tensor,
absf_out: cute.Tensor,
floor_inp: cute.Tensor,
floor_out: cute.Tensor,
):
_unary_ops_kernel(absf_inp, absf_out, floor_inp, floor_out).launch(
grid=[1, 1, 1], block=[absf_inp.shape[0], 1, 1]
@@ -77,7 +81,9 @@ def test_unary_ops():
@cute.kernel
def _binary_ops_kernel(
mag_inp: cute.Tensor, sign_inp: cute.Tensor, out: cute.Tensor,
mag_inp: cute.Tensor,
sign_inp: cute.Tensor,
out: cute.Tensor,
):
tidx, _, _ = cute.arch.thread_idx()
out[tidx] = cute.math.copysign(mag_inp[tidx], sign_inp[tidx])
@@ -85,7 +91,9 @@ def _binary_ops_kernel(
@cute.jit
def _binary_ops_host(
mag_inp: cute.Tensor, sign_inp: cute.Tensor, out: cute.Tensor,
mag_inp: cute.Tensor,
sign_inp: cute.Tensor,
out: cute.Tensor,
):
_binary_ops_kernel(mag_inp, sign_inp, out).launch(
grid=[1, 1, 1], block=[mag_inp.shape[0], 1, 1]