返回 JoyAI-Echo
base_encoder.py
1 import functools
2 from pathlib import Path
3
4 import torch
5 from transformers import AutoImageProcessor, Gemma3ForConditionalGeneration, Gemma3Processor
6
7 from ltx_core.loader.module_ops import ModuleOps
8 from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
9 from ltx_core.utils import find_matching_file
10
11
12 class GemmaTextEncoder(torch.nn.Module):
13 """Pure Gemma text encoder — runs the LLM and returns raw hidden states.
14 The full loader also supports prompt enhancement; the DMD loader retains
15 only the language backbone used by :meth:`encode`.
16 """
17
18 def __init__(
19 self,
20 model: Gemma3ForConditionalGeneration | None = None,
21 tokenizer: LTXVGemmaTokenizer | None = None,
22 processor: Gemma3Processor | None = None,
23 dtype: torch.dtype = torch.bfloat16,
24 ):
25 super().__init__()
26 self.model = model
27 self.tokenizer = tokenizer
28 self.processor = processor
29 self._dtype = dtype
30
31 def encode(
32 self,
33 text: str,
34 padding_side: str = "left",
35 ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
36 """Run Gemma LLM and return raw hidden states + attention mask.
37 Calls the inner model (self.model.model) to skip lm_head logits computation (~500 MiB saving).
38 Returns:
39 (hidden_states, attention_mask) where hidden_states is a tuple of per-layer tensors.
40 """
41 return self.encode_batch([text], padding_side=padding_side)
42
43 def encode_batch(
44 self,
45 texts: list[str],
46 padding_side: str = "left",
47 ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
48 """Run one padded Gemma forward for a non-empty batch of prompts."""
49 if not texts:
50 raise ValueError("texts must contain at least one prompt")
51 if padding_side != "left":
52 raise ValueError("Gemma text encoding requires left padding")
53
54 encoded = self.tokenizer.tokenizer(
55 [text.strip() for text in texts],
56 padding="max_length",
57 max_length=self.tokenizer.max_length,
58 truncation=True,
59 return_tensors="pt",
60 )
61 input_ids = encoded.input_ids.to(self.model.device)
62 attention_mask = encoded.attention_mask.to(self.model.device)
63 outputs = self.model.model(
64 input_ids=input_ids,
65 attention_mask=attention_mask,
66 output_hidden_states=True,
67 use_cache=False,
68 )
69 hidden_states = outputs.hidden_states
70 del outputs
71 return hidden_states, attention_mask
72
73 # --- Prompt enhancement methods ---
74
75 def _enhance(
76 self,
77 messages: list[dict[str, str]],
78 image: torch.Tensor | None = None,
79 max_new_tokens: int = 512,
80 seed: int = 10,
81 ) -> str:
82 text = self.processor.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
83
84 model_inputs = self.processor(
85 text=text,
86 images=image,
87 return_tensors="pt",
88 ).to(self.model.device)
89 pad_token_id = self.processor.tokenizer.pad_token_id if self.processor.tokenizer.pad_token_id is not None else 0
90 model_inputs = _pad_inputs_for_attention_alignment(model_inputs, pad_token_id=pad_token_id)
91
92 with torch.inference_mode(), torch.random.fork_rng(devices=[self.model.device]):
93 torch.manual_seed(seed)
94 outputs = self.model.generate(
95 **model_inputs,
96 max_new_tokens=max_new_tokens,
97 do_sample=True,
98 temperature=0.7,
99 )
100 generated_ids = outputs[0][len(model_inputs.input_ids[0]) :]
101 enhanced_prompt = self.processor.tokenizer.decode(generated_ids, skip_special_tokens=True)
102
103 return enhanced_prompt
104
105 def enhance_t2v(
106 self,
107 prompt: str,
108 max_new_tokens: int = 512,
109 system_prompt: str | None = None,
110 seed: int = 10,
111 ) -> str:
112 """Enhance a text prompt for T2V generation."""
113 system_prompt = system_prompt or self.default_gemma_t2v_system_prompt
114
115 messages = [
116 {"role": "system", "content": system_prompt},
117 {"role": "user", "content": f"user prompt: {prompt}"},
118 ]
119
120 return self._enhance(messages, max_new_tokens=max_new_tokens, seed=seed)
121
122 def enhance_i2v(
123 self,
124 prompt: str,
125 image: torch.Tensor,
126 max_new_tokens: int = 512,
127 system_prompt: str | None = None,
128 seed: int = 10,
129 ) -> str:
130 """Enhance a text prompt for I2V generation using a reference image."""
131 system_prompt = system_prompt or self.default_gemma_i2v_system_prompt
132 messages = [
133 {"role": "system", "content": system_prompt},
134 {
135 "role": "user",
136 "content": [
137 {"type": "image"},
138 {"type": "text", "text": f"User Raw Input Prompt: {prompt}."},
139 ],
140 },
141 ]
142 return self._enhance(messages, image=image, max_new_tokens=max_new_tokens, seed=seed)
143
144 @functools.cached_property
145 def default_gemma_i2v_system_prompt(self) -> str:
146 return _load_system_prompt("gemma_i2v_system_prompt.txt")
147
148 @functools.cached_property
149 def default_gemma_t2v_system_prompt(self) -> str:
150 return _load_system_prompt("gemma_t2v_system_prompt.txt")
151
152
153 # --- Standalone utility functions ---
154
155
156 @functools.lru_cache(maxsize=2)
157 def _load_system_prompt(prompt_name: str) -> str:
158 with open(Path(__file__).parent / "prompts" / f"{prompt_name}", "r") as f:
159 return f.read()
160
161
162 def _cat_with_padding(
163 tensor: torch.Tensor,
164 padding_length: int,
165 value: int | float,
166 ) -> torch.Tensor:
167 """Concatenate a tensor with a padding tensor of the given value."""
168 return torch.cat(
169 [
170 tensor,
171 torch.full(
172 (1, padding_length),
173 value,
174 dtype=tensor.dtype,
175 device=tensor.device,
176 ),
177 ],
178 dim=1,
179 )
180
181
182 def _pad_inputs_for_attention_alignment(
183 model_inputs: dict[str, torch.Tensor],
184 pad_token_id: int = 0,
185 alignment: int = 8,
186 ) -> dict[str, torch.Tensor]:
187 """Pad sequence length to multiple of alignment for Flash Attention compatibility."""
188 seq_len = model_inputs.input_ids.shape[1]
189 padded_len = ((seq_len + alignment - 1) // alignment) * alignment
190 padding_length = padded_len - seq_len
191
192 if padding_length > 0:
193 model_inputs["input_ids"] = _cat_with_padding(model_inputs.input_ids, padding_length, pad_token_id)
194 model_inputs["attention_mask"] = _cat_with_padding(model_inputs.attention_mask, padding_length, 0)
195 if "token_type_ids" in model_inputs and model_inputs["token_type_ids"] is not None:
196 model_inputs["token_type_ids"] = _cat_with_padding(model_inputs["token_type_ids"], padding_length, 0)
197
198 return model_inputs
199
200
201 def tokenizer_module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]:
202 tokenizer_root = str(find_matching_file(gemma_root, "tokenizer.model").parent)
203
204 def load_tokenizer(module: GemmaTextEncoder) -> GemmaTextEncoder:
205 module.tokenizer = LTXVGemmaTokenizer(tokenizer_root, 1024)
206 return module
207
208 tokenizer_load_ops = ModuleOps(
209 "TokenizerLoad",
210 matcher=lambda module: isinstance(module, GemmaTextEncoder) and module.tokenizer is None,
211 mutator=load_tokenizer,
212 )
213 return (tokenizer_load_ops,)
214
215
216 def module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]:
217 processor_root = str(find_matching_file(gemma_root, "preprocessor_config.json").parent)
218
219 def load_processor(module: GemmaTextEncoder) -> GemmaTextEncoder:
220 image_processor = AutoImageProcessor.from_pretrained(processor_root, local_files_only=True)
221 if not module.tokenizer:
222 raise ValueError("Tokenizer model operation must be performed before processor model operation")
223 module.processor = Gemma3Processor(image_processor=image_processor, tokenizer=module.tokenizer.tokenizer)
224 return module
225
226 processor_load_ops = ModuleOps(
227 "ProcessorLoad",
228 matcher=lambda module: isinstance(module, GemmaTextEncoder) and module.processor is None,
229 mutator=load_processor,
230 )
231 return (*tokenizer_module_ops_from_gemma_root(gemma_root), processor_load_ops)
232
232 lines PYTHON