返回 JoyAI-Echo
encoder_configurator.py
根目录 / 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 EMBEDDINGS_PROCESSOR_KEY_OPS = (
125 SDOps("EMBEDDINGS_PROCESSOR_KEY_OPS")
126 # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
127 .with_matching(prefix="text_embedding_projection.aggregate_embed.")
128 .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
129 # V2 dual aggregate embeds
130 .with_matching(prefix="text_embedding_projection.video_aggregate_embed.")
131 .with_replacement("text_embedding_projection.video_aggregate_embed.", "feature_extractor.video_aggregate_embed.")
132 .with_matching(prefix="text_embedding_projection.audio_aggregate_embed.")
133 .with_replacement("text_embedding_projection.audio_aggregate_embed.", "feature_extractor.audio_aggregate_embed.")
134 # 2. Map the connectors
135 .with_matching(prefix="model.diffusion_model.video_embeddings_connector.")
136 .with_replacement("model.diffusion_model.video_embeddings_connector.", "video_connector.")
137 .with_matching(prefix="model.diffusion_model.audio_embeddings_connector.")
138 .with_replacement("model.diffusion_model.audio_embeddings_connector.", "audio_connector.")
139 )
140
141 VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
142 SDOps("VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS")
143 # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
144 .with_matching(prefix="text_embedding_projection.aggregate_embed.")
145 .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
146 # V2 video aggregate embed
147 .with_matching(prefix="text_embedding_projection.video_aggregate_embed.")
148 .with_replacement("text_embedding_projection.video_aggregate_embed.", "feature_extractor.video_aggregate_embed.")
149 # 2. Map the connectors
150 .with_matching(prefix="model.diffusion_model.embeddings_connector.")
151 .with_replacement("model.diffusion_model.embeddings_connector.", "embeddings_processor.video_connector.")
152 )
153
154
155 def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
156 model = module.model
157 v_model = model.model.vision_tower.vision_model
158 l_model = model.model.language_model
159
160 config = model.config.text_config
161 dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
162 base = config.rope_local_base_freq
163 local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
164 inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
165
166 positions_length = len(v_model.embeddings.position_ids[0])
167 position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
168 v_model.embeddings.register_buffer("position_ids", position_ids)
169 embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
170 l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
171 l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
172 l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
173
174 return module
175
176
177 GEMMA_MODEL_OPS = ModuleOps(
178 name="GemmaModel",
179 matcher=lambda module: hasattr(module, "model") and isinstance(module.model, Gemma3ForConditionalGeneration),
180 mutator=create_and_populate,
181 )
182
182 lines PYTHON