Files
revids/app/text_encoder.py

211 lines
8.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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