From 15b37e129bbd1e9e94cbe6debc5d8a2b2f4df338 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 18 Jul 2026 06:24:40 +0000 Subject: [PATCH] Inject hosted pre-cast text encoders during pipeline assembly Wire te_prequant_pipe_kwargs into the three pipeline assembly sites: the diffusion full-pipeline branch, the diffusion transformer-only and GGUF branch (where the companion TE is the big remaining download), and the shared video assembly path before the pipeline/component split. Injection is gated exactly like the runtime cast (mode normalized to fp8, device supported, family not denied), so it can never engage where quantize_text_encoders would not; the later quantize_text_encoders call re-applies the cast idempotently and keeps status reporting truthful. With no hosted checkpoint configured the call returns {} and assembly loads the dense encoder as before. --- studio/backend/core/inference/diffusion.py | 28 ++++++++++++++++++++++ studio/backend/core/inference/video.py | 16 +++++++++++++ 2 files changed, 44 insertions(+) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 3a59ba5dd4..643cfdc171 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -95,6 +95,7 @@ from .diffusion_cache import ( normalize_transformer_cache, ) from .diffusion_precision import TE_QUANT_AUTO, normalize_te_quant, quantize_text_encoders +from .diffusion_te_prequant import te_prequant_pipe_kwargs from .diffusion_vae_quant import VAE_QUANT_AUTO, normalize_vae_quant, quantize_vae from .diffusion_prequant import ( load_prequantized_transformer, @@ -1462,6 +1463,20 @@ class DiffusionBackend: # The repo names a Llama text_encoder_4 it does not ship; # supply it from the open mirror (diffusion_hidream.py). pipe_kwargs.update(hidream_te4_kwargs(dtype, hf_token)) + # A hosted pre-cast fp8 text encoder (when the family ships one and + # the runtime cast would engage) skips the dense TE download; the + # later quantize_text_encoders re-applies the cast idempotently. + pipe_kwargs.update( + te_prequant_pipe_kwargs( + fam, + repo_id, + te_quant_mode = text_encoder_quant, + target = target, + dtype = dtype, + hf_token = hf_token, + logger = logger, + ) + ) # The prefetched snapshot dir keeps from_pretrained off the hub (its # sweep re-pulls files the scoped prefetch skipped: 24 GB per FLUX.1). pipe = pipeline_cls.from_pretrained( @@ -1505,6 +1520,19 @@ class DiffusionBackend: if fam.name == HIDREAM_FAMILY_NAME: # Same Llama TE4 assembly as the full-pipeline branch above. pipe_kwargs.update(hidream_te4_kwargs(dtype, hf_token)) + # Same pre-cast TE injection as the full-pipeline branch: the GGUF + # supplies the transformer, so the companion TE is the big download. + pipe_kwargs.update( + te_prequant_pipe_kwargs( + fam, + base, + te_quant_mode = text_encoder_quant, + target = target, + dtype = dtype, + hf_token = hf_token, + logger = logger, + ) + ) pipe = pipeline_cls.from_pretrained( _base_local_dir or base, **pipe_kwargs ) diff --git a/studio/backend/core/inference/video.py b/studio/backend/core/inference/video.py index 5585d46dcc..4b1053a803 100644 --- a/studio/backend/core/inference/video.py +++ b/studio/backend/core/inference/video.py @@ -1376,6 +1376,22 @@ class VideoBackend: pipe_kwargs["torch_dtype"] = {"vae": torch.float32, "default": dtype} if hf_token: pipe_kwargs["token"] = hf_token + # A hosted pre-cast fp8 text encoder (when the family ships one and the runtime cast + # would engage) skips the dense TE download -- for LTX's Gemma3-27B that is the ~50 GB + # heavyweight of the load. quantize_text_encoders below re-applies the cast idempotently. + from .diffusion_te_prequant import te_prequant_pipe_kwargs + + pipe_kwargs.update( + te_prequant_pipe_kwargs( + fam, + repo_id if kind == "pipeline" else base, + te_quant_mode = text_encoder_quant, + target = target, + dtype = dtype, + hf_token = hf_token, + logger = logger, + ) + ) if kind == "pipeline": # The pre-downloaded snapshot dir keeps from_pretrained off the hub (its sweep would # also pull root checkpoints + duplicate shards); hub id when pre-download was skipped.