LittlehamsterXu opened a new issue, #20315:
URL: https://github.com/apache/tvm/issues/20315

   ## Summary
   
   On targets without native FP8 support, `tirx.transform.FP8ComputeLegalize` 
uses `tvm::tirx::DTypeConversion` for FP8 promotion and narrowing. 
`DTypeConversion` mis-encodes valid `float8_e5m2` zero and subnormal values 
when converting to/from `float32`.
   
   ## Environment
   
   - Target: `llvm` on x86_64 Linux
   - Dependencies: `numpy`, `ml_dtypes`
   
   The runtime reproductions used TVM `v0.25.0.post1` 
(`b3e249b7d75f8f3bc7cbee48188d3c80ae323437`) and a source build based on 
upstream `main` commit `141a187f35e9b6263db6665bc99eb03146d9b23f`. On 
2026-09-11, upstream `main` was `8f328e802cfe5e41fcc8f5c17e7582b1c28bfce4`, 15 
commits ahead; the current source was inspected at that commit, and none of 
those intervening commits modified `dtype_conversion.cc` or 
`unsupported_dtype_legalize.cc`. Both runtime reproductions produced the same 
result below.
   
   ## Reproduction
   
   Save the following as `reproduce-tvm-fp8-zero-subnormal-float32.py` and run 
it in an environment containing TVM `0.25.0.post1`:
   
   ```bash
   python reproduce-tvm-fp8-zero-subnormal-float32.py
   ```
   
   ```python
   import ml_dtypes
   import numpy as np
   import tvm
   from tvm import tirx
   from tvm.script import tirx as T
   
   
   # +0, -0, smallest subnormal, middle subnormal, largest subnormal, smallest 
normal.
   FP8_TO_F32_INPUT_BITS = np.array([0x00, 0x80, 0x01, 0x02, 0x03, 0x04], 
dtype=np.uint8)
   FP8_TO_F32_EXPECTED_BITS = [0x00000000, 0x80000000, 0x37800000, 0x38000000, 
0x38400000, 0x38800000]
   F32_TO_FP8_INPUT_BITS = np.array(
       [0x00000000, 0x80000000, 0x37800000, 0x38000000, 0x38400000, 0x38800000],
       dtype=np.uint32,
   )
   F32_TO_FP8_EXPECTED_BITS = [0x00, 0x80, 0x01, 0x02, 0x03, 0x04]
   
   
   @T.prim_func
   def fp8_to_f32(A: T.Buffer((6,), "float8_e5m2"), B: T.Buffer((6,), 
"float32")):
       for i in T.serial(6):
           B[i] = A[i]
   
   
   @T.prim_func
   def f32_to_fp8(A: T.Buffer((6,), "float32"), B: T.Buffer((6,), 
"float8_e5m2")):
       for i in T.serial(6):
           B[i] = A[i]
   
   
   fp8_to_f32_module = tirx.build(fp8_to_f32, target="llvm")
   fp8_to_f32_src = 
tvm.runtime.tensor(FP8_TO_F32_INPUT_BITS.view(ml_dtypes.float8_e5m2))
   fp8_to_f32_dst = tvm.runtime.tensor(np.zeros(6, dtype=np.float32))
   fp8_to_f32_module["main"](fp8_to_f32_src, fp8_to_f32_dst)
   fp8_to_f32_observed = [int(value) for value in 
fp8_to_f32_dst.numpy().view(np.uint32)]
   
   f32_to_fp8_module = tirx.build(f32_to_fp8, target="llvm")
   f32_to_fp8_src = tvm.runtime.tensor(F32_TO_FP8_INPUT_BITS.view(np.float32))
   f32_to_fp8_dst = tvm.runtime.tensor(np.zeros(6, dtype=ml_dtypes.float8_e5m2))
   f32_to_fp8_module["main"](f32_to_fp8_src, f32_to_fp8_dst)
   f32_to_fp8_observed = [int(value) for value in 
f32_to_fp8_dst.numpy().view(np.uint8)]
   
   print("float8_e5m2 -> float32 observed:", [f"0x{x:08X}" for x in 
fp8_to_f32_observed])
   print("float8_e5m2 -> float32 expected:", [f"0x{x:08X}" for x in 
FP8_TO_F32_EXPECTED_BITS])
   print("float32 -> float8_e5m2 observed:", [f"0x{x:02X}" for x in 
f32_to_fp8_observed])
   print("float32 -> float8_e5m2 expected:", [f"0x{x:02X}" for x in 
F32_TO_FP8_EXPECTED_BITS])
   ```
   
   ## Expected behavior
   
   The six inputs are, in order, `+0`, `-0`, the smallest subnormal, the middle 
subnormal, the largest subnormal, and the smallest normal:
   
   For `float8_e5m2` (`bias = 15`, `mantissa bits = 2`), an exponent field of 
zero denotes signed zero when the mantissa is zero and a subnormal otherwise. 
For a subnormal with mantissa `M > 0`, the value is `(-1)^S × M × 2^-16`; 
normal values use `(-1)^S × (1 + M / 4) × 2^(E - 15)`. Therefore, `0x00` and 
`0x80` are `+0` and `-0`, while `0x01`, `0x02`, and `0x03` are `2^-16`, 
`2^-15`, and `3 × 2^-16`.
   
   | value | `+0` | `-0` | smallest subnormal | middle subnormal | largest 
subnormal | smallest normal |
   | --- | --- | --- | --- | --- | --- | --- |
   | `float8_e5m2` bits | `00` | `80` | `01` | `02` | `03` | `04` |
   | `float32` bits | `00000000` | `80000000` | `37800000` | `38000000` | 
`38400000` | `38800000` |
   
   `float8_e5m2 -> float32` should produce the listed `float32` words, and 
`float32 -> float8_e5m2` should produce the listed FP8 bytes.
   
   ## Actual behavior
   
   The reproducer prints:
   
   ```text
   float8_e5m2 -> float32 observed: 0x38000000 0xB8000000 0x38200000 0x38400000 
0x38600000 0x38800000
   float8_e5m2 -> float32 expected: 0x00000000 0x80000000 0x37800000 0x38000000 
0x38400000 0x38800000
   float32 -> float8_e5m2 observed: 0x00 0x00 0x00 0x00 0x02 0x04
   float32 -> float8_e5m2 expected: 0x00 0x80 0x01 0x02 0x03 0x04
   ```
   
   The smallest-normal control (`0x04`, `2^-14`) is correct in both directions, 
while all tested exponent-zero cases are mis-encoded.
   
   The same output was observed with the rebuilt main-derived source checkout 
described above.
   
   ## Suspected cause
   
   The relevant code is 
[`src/tirx/transform/dtype_conversion.cc`](https://github.com/apache/tvm/blob/b3e249b7d75f8f3bc7cbee48188d3c80ae323437/src/tirx/transform/dtype_conversion.cc#L63-L98).
 The two relevant branches are:
   
   ```cpp
   PrimExpr exponent_before_delta = ((src_uint_value << 1) >> (src_fp.mantissa 
+ 1));
   if (bias_delta >= 0) {
     PrimExpr ret_exponent =
         (bias_delta > 0) ? (cast(tgt_uint, exponent_before_delta + bias_delta) 
<< tgt_fp.mantissa)
                          : (cast(tgt_uint, exponent_before_delta) << 
tgt_fp.mantissa);
     return reinterpret(tgt_dtype, ret_mantissa | ret_exponent | ret_sign);
   } else {
     PrimExpr round_to_zero = exponent_before_delta < (-bias_delta);
     PrimExpr ret_exponent = cast(tgt_uint, exponent_before_delta - 
(-bias_delta))
                             << tgt_fp.mantissa;
     return reinterpret(tgt_dtype, if_then_else(
         round_to_zero, MakeConst(tgt_uint, 0), ret_mantissa | ret_exponent | 
ret_sign));
   }
   ```
   
   - Widening: when the source exponent is `0`, the code still applies the 
normal-value bias delta (`112` for `float8_e5m2 -> float32`). For this input, 
`exponent_before_delta` is `0`, so zero becomes `2^-15` instead of zero.
   - Narrowing: when `exponent_before_delta < -bias_delta`, the underflow 
branch directly returns zero and does not construct a destination FP8 
subnormal. It also loses the sign of `-0.0`.
   
   TVM declares `float8_e5m2` as an IEEE-style format in 
[`dtype_conversion.h`](https://github.com/apache/tvm/blob/b3e249b7d75f8f3bc7cbee48188d3c80ae323437/src/tirx/transform/dtype_conversion.h#L140-L148).
 The smallest-normal control above is correct; the exponent-zero cases are not 
handled. `FP8StorageLegalize` only remaps the storage representation and is not 
implicated in this numeric conversion.
   
   ## Impact
   
   This is a silent numerical correctness issue in the software FP8 
legalization path. A regular `float8_e5m2` zero value is sufficient to trigger 
the mismatch.
   
   This affects CPU builds and any target that falls back to the software FP8 
conversion, including FP8 dequantization at storage/compute boundaries.
   
   ## Suggested regression test
   
   Regression coverage for `+0`, `-0`, subnormals, and the smallest normal in 
both directions would catch this case.
   


-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to