topabaem's picture
Keep VAE parameters out of inference mode when moving devices
44f12a2 verified
Raw
History Blame Contribute Delete
3.62 kB
"""A loaded packed encoder that answers prompts one at a time.
`ltx_conditioning_dump.py` is the batch form of this: it walks a prompt file and
writes a bundle, which is what comparing encoder builds needs. A served
interface needs the model to stay loaded and answer as requests arrive, so this
holds it and exposes `encode`.
The encoding itself is not reimplemented. `hidden_stack_comfy` is imported from
the dump script precisely because its padding behaviour is subtle - ComfyUI
left-pads to 1024 and takes positions over the padded sequence, so the same
prompt lands on different rotary angles depending on how it was fed - and a
second copy of that reasoning would eventually drift from the first.
"""
from __future__ import annotations
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from ltx_conditioning_dump import hidden_stack_comfy
from ltx_prompt_embedding_gate import apply_aggregates, check_aggregates_present
class PackedPromptEncoder:
"""The packed Gemma4-12B encoder plus the LTX dual-linear head.
`resident` keeps the weights packed in VRAM and decodes inside `forward`,
which is what puts the encoder inside 16 GB. It produces conditioning that
is `torch.equal` to the dequantized path, so the choice costs speed and
never quality.
"""
def __init__(self, model_dir: Path, packed_checkpoint: Path,
device: str = "cuda", resident: bool = True):
import torch
from transformers import AutoTokenizer
from ltx_packed_codec import load_packed_model
self.model_dir = Path(model_dir)
self.checkpoint = Path(packed_checkpoint)
self.device = device
# Cheap header read, so a checkpoint missing the aggregate tables fails
# here rather than after an 8 GB load.
check_aggregates_present(self.checkpoint)
self.model = load_packed_model(
str(self.model_dir), self.checkpoint, resident=resident, device=device)
self.model.eval()
self.tokenizer = AutoTokenizer.from_pretrained(str(self.model_dir))
self._torch = torch
def to(self, device: str):
"""Move the encoder between devices and keep `self.device` honest.
Written for hosts that hand out a GPU per request and take it back
afterwards: the weights live in CPU RAM between calls and are walked
over for the encode. `hidden_stack_comfy` builds its input ids on
`self.device`, so setting the attribute is not bookkeeping - forgetting
it produces a device mismatch inside the first embedding lookup.
"""
self.model.to(device)
self.device = device
return self
def encode(self, prompt: str):
"""One prompt as the `[T, 6144]` block the DiT consumes.
6144 is `cross_attention_dim + audio_cross_attention_dim`, and the order
is video then audio. Getting that order backwards still renders - it
renders mud - so the concatenation lives here next to the code that
produced the halves, not in whatever calls this.
"""
torch = self._torch
stacks = hidden_stack_comfy(self.model, self.tokenizer, [prompt], self.device)
branch = apply_aggregates(stacks, self.checkpoint, {})[0]
merged = torch.cat((branch["video"], branch["audio"]), dim=-1)
if merged.shape[-1] != 6144:
raise SystemExit(f"expected a 6144-wide conditioning, got {merged.shape}")
return merged.to(torch.float32).contiguous()