Skip to content

[Fix][Relax] Dequantize fp8 as a float, not an integer - #20508

Open
andrewfish0 wants to merge 2 commits into
apache:mainfrom
andrewfish0:fix-dequantize-fp8
Open

andrewfish0 wants to merge 2 commits into
apache:mainfrom
andrewfish0:fix-dequantize-fp8

Conversation

@andrewfish0

Copy link
Copy Markdown

relax.dequantize truncates fp8 inputs to integers before the scale is applied.

dequantize_compute picks its intermediate dtype with matches_code(FLOAT, BFLOAT) — DLPack codes 2 and 4. Every narrow float carries its own code instead (Float8E3M4 at 7 through Float4E2M1FN at 17), so an fp8 input takes the int32 branch:

# before
dequantized[v_i0, v_i1] = T.Cast("float32", T.Cast("float16", T.Cast("int32", data[v_i0, v_i1])) - zp[v_i0]) * scale[v_i0]

# after
dequantized[v_i0, v_i1] = (T.Cast("float32", data[v_i0, v_i1]) - T.Cast("float32", zp[v_i0])) * scale[v_i0]

A stored -0.9375 becomes 0.

  • Reachable through the public API: relax.quantize already accepts float8_e4m3fn and float8_e5m2 as output dtypes.
  • Easy to miss: the truncation is a no-op wherever the stored value is already integral, which at E4M3's spacing covers most of the top of the range.
  • quantize is unaffected — it decides the same question with "int" in out_dtype, which excludes the fp8 dtypes correctly.

Change

  • Add PrimType.is_float, forwarding to the is_float that tvm_ffi already provides, so the float code list keeps a single home and a future narrow float needs no edit here.
  • Use it in _dequantize, and drop the now-unused DataTypeCode import.

Tests

  • test_dequantize_float8_e4m3fn_to_fp32 in test_transform_legalize_ops_qdq.py. Every existing dequantize case in that file starts from int8, and the fp8 coverage in test_op_qdq.py stops at type inference, so nothing checked the TIR generated for an fp8 input.
  • Confirmed it fails without the fix (structural mismatch on the int32 cast) and passes with it. 277 passed, 3 skipped across test_transform_legalize_ops*.py.

Not covered

  • The same matches_code(FLOAT, BFLOAT) pattern appears at six other Python sites, including legalize_ops/common.py:83, where an fp8 scalar constant would fall through both branches. I have not assessed whether narrow floats reach any of them, and have left them alone.

dequantize_compute chose its intermediate dtype with
matches_code(FLOAT, BFLOAT). Those are DLPack codes 2 and 4, while every
narrow float carries its own code instead -- Float8E3M4 at 7 through
Float4E2M1FN at 17. An fp8 input therefore took the int32 branch and was
truncated to an integer before the zero point was subtracted and the
scale applied, so a stored -0.9375 became 0.

This is reachable through the public API: relax.quantize already accepts
float8_e4m3fn and float8_e5m2 as output dtypes. It stays hidden wherever
the stored value is already integral, which at E4M3's spacing covers
most of the top of the range.

Ask the dtype whether it is a float instead of enumerating codes, via a
PrimType.is_float that forwards to the one tvm_ffi already provides. The
float code list keeps a single home, so a future narrow float needs no
edit here.

quantize is unaffected -- it decides the same question with
"int" in out_dtype, which excludes the fp8 dtypes correctly.
Every dequantize case in this file starts from int8, and the fp8
coverage in test_op_qdq.py stops at type inference, so nothing checked
the TIR generated for an fp8 input.
@andrewfish0

Copy link
Copy Markdown
Author

@tqchen @hunghsiangwang can you review?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant