[Fix][Relax] Dequantize fp8 as a float, not an integer - #20508
Open
andrewfish0 wants to merge 2 commits into
Open
andrewfish0 wants to merge 2 commits into
andrewfish0 wants to merge 2 commits into
Conversation
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.
Author
|
@tqchen @hunghsiangwang can you review? |
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
relax.dequantizetruncates fp8 inputs to integers before the scale is applied.dequantize_computepicks its intermediate dtype withmatches_code(FLOAT, BFLOAT)— DLPack codes 2 and 4. Every narrow float carries its own code instead (Float8E3M4at 7 throughFloat4E2M1FNat 17), so an fp8 input takes theint32branch:A stored
-0.9375becomes0.relax.quantizealready acceptsfloat8_e4m3fnandfloat8_e5m2as output dtypes.quantizeis unaffected — it decides the same question with"int" in out_dtype, which excludes the fp8 dtypes correctly.Change
PrimType.is_float, forwarding to theis_floatthattvm_ffialready provides, so the float code list keeps a single home and a future narrow float needs no edit here._dequantize, and drop the now-unusedDataTypeCodeimport.Tests
test_dequantize_float8_e4m3fn_to_fp32intest_transform_legalize_ops_qdq.py. Every existing dequantize case in that file starts fromint8, and the fp8 coverage intest_op_qdq.pystops at type inference, so nothing checked the TIR generated for an fp8 input.int32cast) and passes with it. 277 passed, 3 skipped acrosstest_transform_legalize_ops*.py.Not covered
matches_code(FLOAT, BFLOAT)pattern appears at six other Python sites, includinglegalize_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.