Add the pre-cast text-encoder checkpoint builder
Applies the runtime layerwise fp8 storage cast to a model's dense text encoder once and saves the cast state dict with baked metadata (format tag, base_model_id, family, scheme, component, te_class, versions) in the layout diffusion_te_prequant.py validates. Resolves the encoder class from the checkpoint's config.architectures so the recorded te_class matches what the pipeline instantiates. CPU-runnable: the cast touches storage dtypes only.
This commit is contained in:
parent
bbc8c232c3
commit
4e9aab60ae
1 changed files with 117 additions and 0 deletions
117
scripts/build_te_prequant_checkpoint.py
Normal file
117
scripts/build_te_prequant_checkpoint.py
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Build a pre-cast text-encoder checkpoint for the Studio TE prequant path.
|
||||
|
||||
Apply the runtime layerwise-fp8 STORAGE cast (``diffusion_precision._cast_fp8``) to a
|
||||
model's dense text encoder ONCE and save the cast state dict, so the backend can load the
|
||||
~half-size artifact (meta-init + ``load_state_dict(assign=True)``, see
|
||||
``core/inference/diffusion_te_prequant.py``) instead of downloading the full bf16 encoder
|
||||
and casting on every load. The cast is a deterministic storage transform, so the loaded
|
||||
encoder is bit-identical to dense-load-then-cast by construction. CPU-runnable: the cast
|
||||
touches storage dtypes only, no kernels.
|
||||
|
||||
python scripts/build_te_prequant_checkpoint.py \
|
||||
--base Lightricks/LTX-2 --family ltx-2 --component text_encoder \
|
||||
--out outputs/te_prequant/ltx2/text_encoder_fp8.pt
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
BACKEND = Path(__file__).resolve().parent.parent / "studio" / "backend"
|
||||
|
||||
|
||||
def main(argv = None) -> int:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument(
|
||||
"--base", required = True, help = "diffusers base repo (carries the component subfolder)"
|
||||
)
|
||||
p.add_argument("--family", required = True, help = "diffusion/video family name or alias")
|
||||
p.add_argument(
|
||||
"--component",
|
||||
default = "text_encoder",
|
||||
help = "pipeline component attribute (also the repo subfolder)",
|
||||
)
|
||||
p.add_argument("--scheme", default = "fp8", choices = ["fp8"])
|
||||
p.add_argument("--out", required = True, help = "output .pt path for the checkpoint")
|
||||
p.add_argument("--dtype", default = "bfloat16", choices = ["bfloat16"])
|
||||
p.add_argument("--hf-token", default = None)
|
||||
args = p.parse_args(argv)
|
||||
|
||||
sys.path.insert(0, str(BACKEND))
|
||||
import torch
|
||||
import transformers
|
||||
|
||||
from core.inference.diffusion_precision import _cast_fp8
|
||||
from core.inference.diffusion_te_prequant import TE_PREQUANT_FORMAT
|
||||
|
||||
# The family is metadata for forensics; detection lives in different modules per
|
||||
# branch (diffusion_families vs video_families), so resolve best-effort by name.
|
||||
family = args.family.strip().lower()
|
||||
|
||||
print(f"== build TE prequant ({family}/{args.component}/{args.scheme}) ==", flush = True)
|
||||
print(f" loading dense encoder from {args.base} (subfolder={args.component}) ...", flush = True)
|
||||
t0 = time.time()
|
||||
config = transformers.AutoConfig.from_pretrained(
|
||||
args.base, subfolder = args.component, token = args.hf_token
|
||||
)
|
||||
# Prefer the checkpoint's own architecture (what the diffusers pipeline instantiates,
|
||||
# e.g. Gemma3ForConditionalGeneration); AutoModel.from_config would give the bare base
|
||||
# class and record a te_class whose state dict the pipeline cannot use.
|
||||
arch = (getattr(config, "architectures", None) or [None])[0]
|
||||
if arch and hasattr(transformers, arch):
|
||||
encoder_cls_name = arch
|
||||
else:
|
||||
encoder = transformers.AutoModel.from_config(config)
|
||||
encoder_cls_name = type(encoder).__name__
|
||||
del encoder
|
||||
encoder = getattr(transformers, encoder_cls_name).from_pretrained(
|
||||
args.base,
|
||||
subfolder = args.component,
|
||||
torch_dtype = torch.bfloat16,
|
||||
token = args.hf_token,
|
||||
)
|
||||
print(f" casting in place (layerwise {args.scheme}) ...", flush = True)
|
||||
|
||||
class _Target:
|
||||
dtype = torch.bfloat16
|
||||
|
||||
_cast_fp8(encoder, _Target())
|
||||
|
||||
state_dict = {
|
||||
k: (v.detach().to("cpu") if hasattr(v, "detach") else v)
|
||||
for k, v in encoder.state_dict().items()
|
||||
}
|
||||
metadata = {
|
||||
"base_model_id": args.base,
|
||||
"family": family,
|
||||
"scheme": args.scheme,
|
||||
"component": args.component,
|
||||
"te_class": encoder_cls_name,
|
||||
"torch_dtype": args.dtype,
|
||||
"cast_backend": "diffusers_layerwise",
|
||||
"torch_version": torch.__version__,
|
||||
"transformers_version": transformers.__version__,
|
||||
}
|
||||
ckpt = {
|
||||
"format": TE_PREQUANT_FORMAT,
|
||||
"metadata": metadata,
|
||||
"state_dict": state_dict,
|
||||
}
|
||||
out = Path(args.out)
|
||||
out.parent.mkdir(parents = True, exist_ok = True)
|
||||
torch.save(ckpt, out)
|
||||
size_gb = out.stat().st_size / 1e9
|
||||
print(f" saved {out} ({size_gb:.2f} GB) in {time.time() - t0:.0f}s", flush = True)
|
||||
print(f" metadata: {metadata}", flush = True)
|
||||
print("BUILD-TE-PREQUANT-DONE", flush = True)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Loading…
Add table
Add a link
Reference in a new issue