Fix dtype_byte_size reporting zero bytes for sub-byte and packed dtypes - #4237
Open
VihaanAgarwal wants to merge 1 commit into
Open
Fix dtype_byte_size reporting zero bytes for sub-byte and packed dtypes#4237VihaanAgarwal wants to merge 1 commit into
VihaanAgarwal wants to merge 1 commit into
Conversation
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.
What does this PR do?
dtype_byte_sizefalls back to reading the trailing digits of the dtype name:For torch's sub-byte and packed dtypes that trailing number is not the storage size, and the integer division turns it into zero:
compute_module_sizesand thereforeinfer_auto_device_mapthen count such parameters as free, so a device map can overcommit a device by the whole size of those weights.torch has exposed
dtype.itemsizesince 2.1, the same version that already gates the FP8 branch. This PR returnsitemsizeunder that gate, which also replaces the hardcoded FP8 name list. The regex stays as the fallback for torch 2.0.torch.booland theCustomDtypevalues keep their explicit sizes.The
float8_e8m0fnuand sub-byte cases intest_dtype_byte_sizefail onmainwith0 != 1and pass with the change.tests/test_modeling_utils.pypasses locally (torch 2.14, CPU).Before submitting
Who can review?
@SunMarc