返回 JoyAI-Echo
encoder_configurator.py
根目录 / echo_longvideo / ltx-core / src / ltx_core / text_encoders / gemma / encoders / encoder_configurator.py
1 import torch
2 from transformers import Gemma3Config
3 from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS
4 from transformers.models.gemma3 import Gemma3ForConditionalGeneration
5
6 from ltx_core.loader import KeyValueOperationResult
7 from ltx_core.loader.module_ops import ModuleOps
8 from ltx_core.loader.sd_ops import SDOps
9 from ltx_core.model.model_protocol import ModelConfigurator
10 from ltx_core.text_encoders.gemma.config import GEMMA3_CONFIG_FOR_LTX
11 from ltx_core.text_encoders.gemma.embeddings_connector import (
12 AudioEmbeddings1DConnectorConfigurator,
13 Embeddings1DConnectorConfigurator,
14 )
15 from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor
16 from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
17 from ltx_core.text_encoders.gemma.feature_extractor import (
18 FeatureExtractorV1,
19 FeatureExtractorV2,
20 )
21
22
23 class GemmaTextEncoderConfigurator(ModelConfigurator[GemmaTextEncoder]):
24 @classmethod
25 def from_config(cls, config: dict) -> GemmaTextEncoder: # noqa: ARG003
26 gemma_config = Gemma3Config.from_dict(GEMMA3_CONFIG_FOR_LTX.to_dict())
27 with torch.device("meta"):
28 model = Gemma3ForConditionalGeneration(gemma_config)
29
30 return GemmaTextEncoder(model=model)
31
32
33 class EmbeddingsProcessorConfigurator(ModelConfigurator[EmbeddingsProcessor]):
34 @classmethod
35 def from_config(cls, config: dict) -> EmbeddingsProcessor:
36 transformer_config = config.get("transformer", {})
37
38 # Create video embeddings connector (always needed)
39 video_connector = Embeddings1DConnectorConfigurator.from_config(config)
40
41 # Create audio embeddings connector
42 audio_connector = AudioEmbeddings1DConnectorConfigurator.from_config(config)
43
44 # Create feature extractor
45 feature_extractor = _create_feature_extractor(transformer_config)
46
47 return EmbeddingsProcessor(
48 video_connector=video_connector,
49 audio_connector=audio_connector,
50 feature_extractor=feature_extractor,
51 )
52
53
54 _V2_EXPECTED_CONFIG = {
55 "caption_proj_before_connector": True,
56 "caption_projection_first_linear": False,
57 "caption_proj_input_norm": False,
58 "caption_projection_second_linear": False,
59 }
60
61
62 def _create_feature_extractor(transformer_config: dict) -> torch.nn.Module:
63 """Select and create the appropriate feature extractor based on config.
64 Detection logic:
65 - V1: V2 config keys absent → projection lives in transformer
66 - V2: V2 config keys present with exact expected values → per-token RMS norm with dual aggregate embeds
67 - Anything else: NotImplementedError (config drift)
68 """
69 gemma_text_config = GEMMA3_CONFIG_FOR_LTX.text_config
70 embedding_dim = gemma_text_config.hidden_size
71 num_layers = gemma_text_config.num_hidden_layers + 1 # +1 for the embedding layer
72 flat_dim = embedding_dim * num_layers
73
74 overlapping_keys = transformer_config.keys() & _V2_EXPECTED_CONFIG.keys()
75 if not overlapping_keys:
76 aggregate_embed = torch.nn.Linear(flat_dim, embedding_dim, bias=False)
77 return FeatureExtractorV1(aggregate_embed=aggregate_embed, is_av=True)
78
79 missing_keys = _V2_EXPECTED_CONFIG.keys() - overlapping_keys
80 if missing_keys:
81 raise NotImplementedError("Partial V2 config — missing keys: " + ", ".join(sorted(missing_keys)))
82
83 unexpected_value_keys = {k for k in overlapping_keys if transformer_config[k] != _V2_EXPECTED_CONFIG[k]}
84 if unexpected_value_keys:
85 raise NotImplementedError(
86 "Unknown config: "
87 + ", ".join(
88 f"{k}={transformer_config[k]!r} (expected {_V2_EXPECTED_CONFIG[k]!r})" for k in unexpected_value_keys
89 )
90 )
91
92 video_inner_dim = transformer_config["num_attention_heads"] * transformer_config["attention_head_dim"]
93 audio_inner_dim = transformer_config["audio_num_attention_heads"] * transformer_config["audio_attention_head_dim"]
94 return FeatureExtractorV2(
95 video_aggregate_embed=torch.nn.Linear(flat_dim, video_inner_dim, bias=True),
96 embedding_dim=embedding_dim,
97 audio_aggregate_embed=torch.nn.Linear(flat_dim, audio_inner_dim, bias=True),
98 )
99
100
101 # --- Split SDOps: Gemma LLM keys vs Embeddings Processor keys ---
102
103 GEMMA_LLM_KEY_OPS = (
104 SDOps("GEMMA_LLM_KEY_OPS")
105 # 1. Map language model layers (note the double .model prefix)
106 .with_matching(prefix="language_model.model.")
107 .with_replacement("language_model.model.", "model.model.language_model.")
108 # 2. Map the Vision Tower
109 .with_matching(prefix="vision_tower.")
110 .with_replacement("vision_tower.", "model.model.vision_tower.")
111 # 3. Map the Multi-Modal Projector
112 .with_matching(prefix="multi_modal_projector.")
113 .with_replacement("multi_modal_projector.", "model.model.multi_modal_projector.")
114 # 4. Duplicate embed_tokens to lm_head (needed for prompt enhancement via generate())
115 .with_kv_operation(
116 operation=lambda key, value: [
117 KeyValueOperationResult(key, value),
118 KeyValueOperationResult("model.lm_head.weight", value),
119 ],
120 key_prefix="model.model.language_model.embed_tokens.weight",
121 )
122 )
123
124 # DMD inference only consumes hidden states from the language backbone. Filtering
125 # here prevents unused vision/projector/lm-head tensors from being read from the
126 # safetensors shards in the first place.
127 GEMMA_LANGUAGE_ONLY_KEY_OPS = (
128 SDOps("GEMMA_LANGUAGE_ONLY_KEY_OPS")
129 .with_matching(prefix="language_model.model.")
130 .with_replacement("language_model.model.", "model.model.language_model.")
131 )
132
133 EMBEDDINGS_PROCESSOR_KEY_OPS = (
134 SDOps("EMBEDDINGS_PROCESSOR_KEY_OPS")
135 # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
136 .with_matching(prefix="text_embedding_projection.aggregate_embed.")
137 .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
138 # V2 dual aggregate embeds
139 .with_matching(prefix="text_embedding_projection.video_aggregate_embed.")
140 .with_replacement("text_embedding_projection.video_aggregate_embed.", "feature_extractor.video_aggregate_embed.")
141 .with_matching(prefix="text_embedding_projection.audio_aggregate_embed.")
142 .with_replacement("text_embedding_projection.audio_aggregate_embed.", "feature_extractor.audio_aggregate_embed.")
143 # 2. Map the connectors
144 .with_matching(prefix="model.diffusion_model.video_embeddings_connector.")
145 .with_replacement("model.diffusion_model.video_embeddings_connector.", "video_connector.")
146 .with_matching(prefix="model.diffusion_model.audio_embeddings_connector.")
147 .with_replacement("model.diffusion_model.audio_embeddings_connector.", "audio_connector.")
148 )
149
150 VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
151 SDOps("VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS")
152 # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
153 .with_matching(prefix="text_embedding_projection.aggregate_embed.")
154 .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
155 # V2 video aggregate embed
156 .with_matching(prefix="text_embedding_projection.video_aggregate_embed.")
157 .with_replacement("text_embedding_projection.video_aggregate_embed.", "feature_extractor.video_aggregate_embed.")
158 # 2. Map the connectors
159 .with_matching(prefix="model.diffusion_model.embeddings_connector.")
160 .with_replacement("model.diffusion_model.embeddings_connector.", "embeddings_processor.video_connector.")
161 )
162
163
164 def _populate_language_buffers(model: Gemma3ForConditionalGeneration) -> None:
165 l_model = model.model.language_model
166
167 config = model.config.text_config
168 dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
169 base = config.rope_local_base_freq
170 local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
171 inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
172
173 embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
174 l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
175 l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
176 l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
177
178
179 def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
180 model = module.model
181 v_model = model.model.vision_tower.vision_model
182
183 positions_length = len(v_model.embeddings.position_ids[0])
184 position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
185 v_model.embeddings.register_buffer("position_ids", position_ids)
186 _populate_language_buffers(model)
187
188 return module
189
190
191 def create_language_only_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
192 model = module.model
193 _populate_language_buffers(model)
194
195 # Remove meta modules that DMD text encoding never executes. This also keeps
196 # the model builder's uninitialized-parameter check meaningful after the
197 # corresponding checkpoint keys have been filtered out.
198 model.lm_head = None
199 model.model.vision_tower = None
200 model.model.multi_modal_projector = None
201 module.processor = None
202 return module
203
204
205 GEMMA_MODEL_OPS = ModuleOps(
206 name="GemmaModel",
207 matcher=lambda module: hasattr(module, "model") and isinstance(module.model, Gemma3ForConditionalGeneration),
208 mutator=create_and_populate,
209 )
210
211 GEMMA_LANGUAGE_ONLY_MODEL_OPS = ModuleOps(
212 name="GemmaLanguageOnlyModel",
213 matcher=lambda module: hasattr(module, "model") and isinstance(module.model, Gemma3ForConditionalGeneration),
214 mutator=create_language_only_and_populate,
215 )
216
216 lines PYTHON