Move NF4 Gemma encoder into app/, patch ltx_pipelines at runtime

This commit is contained in:
2026-06-03 19:14:47 -04:00
parent ab580a2004
commit 188aaa44a1
2 changed files with 266 additions and 0 deletions
+210
View File
@@ -0,0 +1,210 @@
"""NF4 quantized Gemma prompt encoder owned by this project, patches ltx_pipelines at runtime."""
from __future__ import annotations
import logging
from collections.abc import Iterator
from contextlib import AbstractContextManager, contextmanager
import torch
from ltx_core.block_streaming import StreamingModelBuilder
from ltx_core.loader.primitives import BuilderProtocol
from ltx_core.loader.registry import DummyRegistry, Registry
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.text_encoders.gemma import (
EMBEDDINGS_PROCESSOR_KEY_OPS,
GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS,
EmbeddingsProcessorConfigurator,
GemmaTextEncoderConfigurator,
GemmaTextEncoder,
module_ops_from_gemma_root,
)
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor, EmbeddingsProcessorOutput
from ltx_core.text_encoders.gemma.tokenizer import GemmaTokenizer as LTXVGemmaTokenizer
from ltx_core.utils import find_matching_file
from ltx_pipelines.utils.blocks import PromptEncoder as BasePromptEncoder, _streaming_model
from ltx_pipelines.utils.gpu_model import gpu_model
from ltx_pipelines.utils.helpers import generate_enhanced_prompt
from ltx_pipelines.utils.types import OffloadMode
logger = logging.getLogger(__name__)
class Nf4PromptEncoder(BasePromptEncoder):
"""PromptEncoder with NF4 quantized Gemma support."""
def __init__(
self,
checkpoint_path: str,
gemma_root: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
offload_mode: OffloadMode = OffloadMode.NONE,
text_encoder_builder: BuilderProtocol | None = None,
gemma_quantization: str | None = None,
) -> None:
self._gemma_root = gemma_root
self._checkpoint_path = checkpoint_path
self._dtype = dtype
self._device = device
self._offload_mode = offload_mode
self._gemma_quantization = gemma_quantization
self._cached_quantized_encoder: GemmaTextEncoder | None = None
self._processor = None # type: ignore
if gemma_quantization == "nf4":
self._is_quantized = True
elif gemma_quantization is not None and gemma_quantization != "":
logger.warning(
"Unsupported gemma_quantization '%s', falling back to bf16", gemma_quantization
)
self._is_quantized = False
else:
self._is_quantized = False
if text_encoder_builder is not None:
if offload_mode != OffloadMode.NONE:
raise ValueError(
"text_encoder_builder cannot be used with offload_mode != OffloadMode.NONE"
)
self._text_encoder_builder = text_encoder_builder
self._streaming_text_encoder_builder = None
elif not self._is_quantized:
module_ops = module_ops_from_gemma_root(gemma_root)
model_folder = find_matching_file(gemma_root, "model*.safetensors").parent
weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")]
self._text_encoder_builder = Builder(
model_path=tuple(weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=GEMMA_LLM_KEY_OPS,
module_ops=(GEMMA_MODEL_OPS, *module_ops),
registry=registry or DummyRegistry(),
)
self._streaming_text_encoder_builder = StreamingModelBuilder(
model_path=tuple(weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=GEMMA_LLM_KEY_OPS,
module_ops=(GEMMA_MODEL_OPS, *module_ops),
registry=registry or DummyRegistry(),
blocks_attr="model.model.language_model.layers",
blocks_prefix="model.model.language_model.layers",
)
self._embeddings_processor_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=EmbeddingsProcessorConfigurator,
model_sd_ops=EMBEDDINGS_PROCESSOR_KEY_OPS,
registry=registry or DummyRegistry(),
)
@staticmethod
def _try_import_bitsandbytes() -> None:
try:
import bitsandbytes # noqa: F401
except ImportError as e:
raise RuntimeError(
"NF4 Gemma quantization requires bitsandbytes. Install with:\n"
" pip install bitsandbytes\n"
"Or disable REVIDS_LTX_GEMMA_QUANTIZATION to use bf16."
) from e
def _load_quantized_gemma(self) -> GemmaTextEncoder:
self._try_import_bitsandbytes()
from transformers import (
AutoImageProcessor,
BitsAndBytesConfig,
Gemma3ForConditionalGeneration,
Gemma3Processor,
)
gemma_path = str(find_matching_file(self._gemma_root, "model*.safetensors").parent)
tokenizer_path = str(find_matching_file(self._gemma_root, "tokenizer.model").parent)
processor_path = str(find_matching_file(self._gemma_root, "preprocessor_config.json").parent)
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
llm_int8_skip_modules=["lm_head"],
)
with torch.device("meta"):
gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
gemma_path,
quantization_config=bnb_config,
torch_dtype=torch.bfloat16,
device_map={"": self._device},
local_files_only=True,
)
tokenizer = LTXVGemmaTokenizer(tokenizer_path, 1024)
image_processor = AutoImageProcessor.from_pretrained(
processor_path, local_files_only=True, use_fast=False
)
self._processor = Gemma3Processor(
image_processor=image_processor, tokenizer=tokenizer.tokenizer
)
return GemmaTextEncoder(
model=gemma_model, tokenizer=tokenizer, processor=self._processor, dtype=self._dtype
)
def _build_text_encoder(self) -> torch.nn.Module:
if self._is_quantized:
if self._cached_quantized_encoder is None:
logger.info("Loading NF4 quantized Gemma encoder from %s", self._gemma_root)
self._cached_quantized_encoder = self._load_quantized_gemma()
logger.info("NF4 Gemma encoder loaded and cached")
return self._cached_quantized_encoder.eval()
return self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval()
def _build_embeddings_processor(self) -> EmbeddingsProcessor:
return self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).eval()
def _text_encoder_ctx(self) -> AbstractContextManager:
if self._is_quantized:
@contextmanager
def _cached_ctx() -> Iterator:
yield self._build_text_encoder()
return _cached_ctx()
if self._offload_mode != OffloadMode.NONE:
return _streaming_model(
self._streaming_text_encoder_builder,
self._offload_mode,
self._device,
self._dtype,
)
return gpu_model(self._build_text_encoder())
def __call__(
self,
prompts: list[str],
*,
enhance_first_prompt: bool = False,
enhance_prompt_image: str | None = None,
enhance_prompt_seed: int = 42,
) -> list[EmbeddingsProcessorOutput]:
logger.info("Building text encoder from %s", self._gemma_root)
with self._text_encoder_ctx() as text_encoder: # type: ignore[var-annotated]
if enhance_first_prompt:
prompts = list(prompts)
prompts[0] = generate_enhanced_prompt(
text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed
)
raw_outputs = [text_encoder.encode(p) for p in prompts]
logger.info(
"Text encoder done, building embeddings processor from %s", self._checkpoint_path
)
with gpu_model(self._build_embeddings_processor()) as embeddings_processor:
result = [
embeddings_processor.process_hidden_states(hs, mask)
for hs, mask in raw_outputs
]
logger.info("Prompt encoding complete")
return result