Fix pre-cast TE checkpoint loading and engagement reporting
Two bugs found while building the hosted checkpoints: - The builder recorded torch.__version__ (a TorchVersion object) in the checkpoint metadata, so torch.load(weights_only=True) rejected every artifact and the loader silently fell back to the dense download. Record plain strings. - Re-applying the layerwise fp8 cast to an injected pre-cast encoder raised on the duplicate hook registration, making quantize_text_encoders report the engaged cast as failed (status showed no TE quant while the encoder ran fp8). _cast_fp8 now returns early when the hooks are already installed. Also corrects the LTX TE size note: Gemma3-12B stored fp32 (~49 GB), not 27B.
This commit is contained in:
parent
9c422fd50e
commit
5249a52356
3 changed files with 21 additions and 3 deletions
|
|
@ -95,8 +95,10 @@ def main(argv = None) -> int:
|
|||
"te_class": encoder_cls_name,
|
||||
"torch_dtype": args.dtype,
|
||||
"cast_backend": "diffusers_layerwise",
|
||||
"torch_version": torch.__version__,
|
||||
"transformers_version": transformers.__version__,
|
||||
# str(): torch.__version__ is a TorchVersion object; pickling it into the
|
||||
# checkpoint makes torch.load(weights_only=True) reject the whole artifact.
|
||||
"torch_version": str(torch.__version__),
|
||||
"transformers_version": str(transformers.__version__),
|
||||
}
|
||||
ckpt = {
|
||||
"format": TE_PREQUANT_FORMAT,
|
||||
|
|
|
|||
|
|
@ -450,6 +450,12 @@ def _cast_fp8(encoder: Any, target: Any) -> None:
|
|||
from diffusers.hooks import apply_layerwise_casting
|
||||
from diffusers.hooks.layerwise_casting import DEFAULT_SKIP_MODULES_PATTERN
|
||||
|
||||
# Idempotent: a pre-cast encoder (diffusion_te_prequant) arrives with the layerwise hooks
|
||||
# already installed, and re-registering the same hook name raises -- which would make
|
||||
# quantize_text_encoders report the (actually engaged) cast as failed.
|
||||
if _has_layerwise_hooks(encoder):
|
||||
return
|
||||
|
||||
# Layerwise casting stores each leaf's weights in fp8 and upcasts per forward. Two things on a
|
||||
# transformers encoder push an fp8 weight/activation into an op that can't handle it, both
|
||||
# crashing only at generation, so skip the offending modules:
|
||||
|
|
@ -487,6 +493,16 @@ def _cast_fp8(encoder: Any, target: Any) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _has_layerwise_hooks(encoder: Any) -> bool:
|
||||
"""True when any submodule already carries the diffusers layerwise-casting hook."""
|
||||
for module in encoder.modules():
|
||||
registry = getattr(module, "_diffusers_hook", None)
|
||||
get_hook = getattr(registry, "get_hook", None)
|
||||
if callable(get_hook) and get_hook("layerwise_casting") is not None:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _cast_nvfp4(encoder: Any, target: Any) -> None:
|
||||
# Weight-only NVFP4: linear weights become 4-bit NVFP4 on Blackwell FP4 cores; norms /
|
||||
# embeddings untouched. Exclude the VLM vision tower / lm_head / T5 wo and sub-512 projections
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@
|
|||
|
||||
The runtime ``text_encoder_quant=fp8`` path (``diffusion_precision._cast_fp8``) downloads
|
||||
the full bf16 text encoder and layerwise-casts it in place on every load. For the
|
||||
heavyweight encoders (LTX's Gemma3-27B ~50 GB, FLUX.2-dev's Mistral-24B ~48 GB,
|
||||
heavyweight encoders (LTX's Gemma3-12B ~49 GB fp32, FLUX.2-dev's Mistral-24B ~48 GB,
|
||||
Qwen-Image's Qwen2.5-VL ~16.6 GB) that download dominates a fresh machine's load. When the
|
||||
encoder was already cast and saved (``scripts/build_te_prequant_checkpoint.py``), this loads
|
||||
the ~half-size fp8-storage state dict directly: meta-init the encoder skeleton,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue