| 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 |