Handle transposed fp8 weights in LoRA fusion - #293
Draft
masahiroteraoka wants to merge 1 commit into
Draft
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.
PR: Handle transposed fp8 weights during scaled-mm LoRA fusion
Summary
_fp8_scaled_mm_fuseinpackages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.pyassumes that pre-quantized fp8 weights are stored in the standard(out, in)layout. Some fp8 checkpoints store linear weights in transposed(in, out)layout for scaled matrix multiplication.When LoRA fusion is applied, the LoRA delta is produced in standard
(out, in)layout. If the checkpoint weight is transposed, the existing dequantization path tries to add tensors with incompatible shapes and raises a shape mismatch error.Fix
Make scaled-mm fp8 LoRA fusion shape-aware:
weight.shape == deltas.shape, dequantize the weight as-is.weight.t().shape == deltas.shape, transpose before dequantizing.ValueErrorthat includes both shapes.This keeps the existing standard-layout path unchanged while allowing transposed fp8 checkpoint weights to fuse with LoRA deltas correctly.
Compatibility
(out, in)fp8 checkpoints keep the same behavior.(in, out)fp8 checkpoints are now handled before adding the LoRA delta.quantize_weight_to_fp8_per_tensorpath.Testing
python -m py_compile packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py