Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
[AMD][CDNA4] Address review comments for MXFP4 gfx950 support
Four issues from CodeRabbit review fixed:

1. hip_fp4.h: Fix FP4 E2M1 denormal decoder round-trip bug.
   The decoder returned 0.25f for mant==1 (exp==0) but the encoder
   maps nibble 1 to 0.5f, breaking float→fp4→float round-trips.
   Fix: return 0.5f to match the encoder and LUT tables.

2. codegen_hip.cc: Fix UB address-of-prvalue in FP4→float16 cast.
   The generated code took &(__tl_cvt_fp4x2_to_half2(...)), which
   takes the address of a temporary (C++ UB). Fix: materialize the
   uint1 return value into a local variable first, matching the
   pattern already used in the float32/bfloat16 cast paths.

3. codegen_hip.cc: Extend fp4_pair_cast to cover FP4x16 and FP4x32.
   Previously only 2/4/8-lane FP4 casts were handled; x16/x32 fell
   through silently to generic vector code that cannot index fp4_e2_*
   aggregates. The pairwise uint8_t byte logic is correct for any
   even lane count, so extend to include 16 and 32 lanes.

4. mxfp.py: Narrow bare except to expected import-time errors.
   The catch-all except Exception silently swallowed all errors from
   target_is_gfx950(), causing HIP targets to fall back to CUDA/PTX.
   Fix: catch only ImportError/ModuleNotFoundError/AttributeError.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
  • Loading branch information
zhangnju and claude committed May 15, 2026
commit e1c3005aab1b3d2e065aed61a55f586f6fb661b9
25 changes: 15 additions & 10 deletions src/backend/rocm/codegen/codegen_hip.cc
Original file line number Diff line number Diff line change
Expand Up @@ -804,7 +804,11 @@ void CodeGenTileLangHIP::VisitExpr_(const CastNode *op, std::ostream &os) {
// PrintType when a float4 type is encountered).
// ---------------------------------------------------------------------------
int fp4_lanes = from_ty.lanes();
bool fp4_pair_cast = (fp4_lanes == 2 || fp4_lanes == 4 || fp4_lanes == 8);
// Pairwise cast: process 2 FP4 lanes at a time via packed uint8_t byte.
// Supported lane widths: 2, 4, 8, 16, 32 (all even widths up to FP4x32).
bool fp4_pair_cast =
(fp4_lanes == 2 || fp4_lanes == 4 || fp4_lanes == 8 || fp4_lanes == 16 ||
fp4_lanes == 32);

Comment thread
coderabbitai[bot] marked this conversation as resolved.
// FP4 -> float16 : use __tl_cvt_fp4x2_to_half2 per 2-element pair
if (from_ty.is_float4_e2m1fn() && target_ty.is_float16() && fp4_pair_cast) {
Expand All @@ -815,16 +819,17 @@ void CodeGenTileLangHIP::VisitExpr_(const CastNode *op, std::ostream &os) {
std::string src = SSAGetID(PrintExpr(op->value), from_ty);
// Iterate over pairs: src is stored as fp4_e2_{lanes}_t; we access the
// packed byte for each pair via reinterpret as uint8_t array.
// Materialize the uint1 result into a local variable to avoid taking the
// address of a prvalue (C++ UB).
for (int i = 0; i < fp4_lanes; i += 2) {
std::ostringstream val;
val << "__tl_cvt_fp4x2_to_half2(((uint8_t*)&(" << src << "))[" << i / 2
<< "])";
// Store both elements of the half2
std::ostringstream v0, v1;
v0 << "((half_t*)(&(" << val.str() << ")))[0]";
v1 << "((half_t*)(&(" << val.str() << ")))[1]";
PrintVecElemStore(sret, target_ty, i, v0.str());
PrintVecElemStore(sret, target_ty, i + 1, v1.str());
std::string tmp = name_supply_->FreshName("_fp4h2_");
this->PrintIndent();
stream << "uint1 " << tmp << " = __tl_cvt_fp4x2_to_half2(((uint8_t*)&("
<< src << "))[" << i / 2 << "]);\n";
PrintVecElemStore(sret, target_ty, i,
"((half_t*)(&(" + tmp + ")))[0]");
PrintVecElemStore(sret, target_ty, i + 1,
"((half_t*)(&(" + tmp + ")))[1]");
}
os << sret;
return;
Expand Down
4 changes: 2 additions & 2 deletions src/tl_templates/hip/hip_fp4.h
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,8 @@ struct fp4_e2_t {
uint32_t mant = bits & 0x1u;
float result;
if (exp == 0u) {
// Denormal: value = (-1)^s * 2^(-1) * (0 + m*0.5) = (-1)^s * m * 0.25
result = mant ? 0.25f : 0.0f;
// Denormal: value = (-1)^s * 0.5 * m (mant==1 => 0.5, matching encoder nibble 1)
result = mant ? 0.5f : 0.0f;
} else {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
// Normal: value = (-1)^s * 2^(e-1) * (1 + m*0.5)
float mantissa = 1.0f + mant * 0.5f;
Expand Down
3 changes: 2 additions & 1 deletion tilelang/quantize/mxfp.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,7 +203,8 @@ def get_mxfp_intrin_group(
from tilelang.utils.target import target_is_gfx950

_is_gfx950 = target_is_gfx950(target)
except Exception:
except (ImportError, ModuleNotFoundError, AttributeError):
# target_is_gfx950 unavailable in this build; assume non-gfx950.
pass
Comment thread
coderabbitai[bot] marked this conversation as resolved.

dtype_map = {T.float16: "f16", T.bfloat16: "bf16"}
Expand Down
Loading