Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions src/cuda/codegen/codegen_cuda.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1811,6 +1811,19 @@ void CodeGenTileLangCUDA::VisitExpr_(const CastNode *op, std::ostream &os) {
// 2. Add a new if statement like the one below.
// 3. In src/cuda/target_utils.cc, allow this vectorizable cast.

// `round` and its `rbits` operand have a PTX lowering only for an f32 source,
// in the four groups below that consult `cast_round` themselves. Every other
// source dtype reaches a vectorized branch that emits the plain conversion
// helper and returns, so it never reaches the rejection at the end of this
// function and the argument is dropped silently -- while the same cast
// written in a scalar loop is rejected. Refuse it for those sources here.
if (!cast_round.empty() && !(from_ty.is_float() && from_ty.bits() == 32)) {
LOG(FATAL) << "round '" << cast_round << "' is not supported for cast from "
<< from_ty << " to " << target_ty
<< " (only supported f32 packed stochastic conversions are "
"available)";
}

// Handle conversion from float16 to float32
if (from_ty.is_float16() && target_ty.is_float() && target_ty.bits() == 32) {
// Use __half22float2 for vectorized conversion (half2 -> float2)
Expand Down
44 changes: 44 additions & 0 deletions testing/python/language/test_tilelang_cast_rounding.py
Original file line number Diff line number Diff line change
Expand Up @@ -522,5 +522,49 @@ def test_cast_rs_rbits_invariance(rbits_mode, dtype):
assert expected in code, f"Expected '{expected}' for rbits_mode={rbits_mode}, dtype={dtype}.\nGenerated code:\n{code}"


# ===========================================================================
# A non-f32 source cannot honour `round`, in either loop form
# ===========================================================================

_RS_TARGET = {"kind": "cuda", "arch": "sm_90a"}


def _vectorized_f16_to_fp8_kernel(*, round_mode):
@T.prim_func
def main(A: T.Tensor((8,), "float16"), B: T.Tensor((8,), "float8_e4m3fn")):
with T.Kernel(1, threads=1):
for i in T.vectorized(8):
if round_mode is None:
B[i] = cast(A[i], "float8_e4m3fn")
else:
B[i] = cast(A[i], "float8_e4m3fn", round=round_mode, rbits=T.uint32(0x12345678))

return main


@tilelang.testing.requires_cuda
def test_cast_rs_rejected_for_non_f32_source_in_a_vectorized_loop():
"""The vectorized branches must not drop `round` that they cannot honour.

`round` and its `rbits` operand have a PTX lowering only for an f32 source. Every
other source dtype reaches a vectorized branch that emits the plain conversion
helper and returns, before the rejection at the end of the cast lowering is
reached -- so the request was dropped silently: the output was bit-identical to a
plain round-to-nearest cast for every seed, while the same cast written in a
scalar loop was rejected.
"""
with pytest.raises(Exception, match="only supported f32 packed stochastic conversions"), tvm.target.Target(_RS_TARGET):
tilelang.lower(_vectorized_f16_to_fp8_kernel(round_mode="rs"),>


@tilelang.testing.requires_cuda
def test_vectorized_cast_without_round_still_lowers():
"""Control: the same vectorized cast with no rounding mode is unaffected."""
with tvm.target.Target(_RS_TARGET):
artifact = tilelang.lower(_vectorized_f16_to_fp8_kernel(round_mode=None),>

assert "__tl_cvt_half2_to_fp8x2" in artifact.kernel_source, artifact.kernel_source


if __name__ == "__main__":
tilelang.testing.main()
Loading