Skip to content

gemma

Classes

fastvideo.models.encoders.gemma.Embeddings1DConnector

Embeddings1DConnector(config: GemmaConnectorConfig)

Bases: Module

Transformer connector that refines Gemma embeddings for LTX-2.

Source code in fastvideo/models/encoders/gemma.py
def __init__(self, config: GemmaConnectorConfig) -> None:
    super().__init__()
    self.num_attention_heads = config.num_attention_heads
    self.inner_dim = config.num_attention_heads * config.attention_head_dim
    self.positional_embedding_theta = config.positional_embedding_theta
    self.positional_embedding_max_pos = config.positional_embedding_max_pos
    self.rope_type = config.rope_type
    self.double_precision_rope = config.double_precision_rope
    self.transformer_1d_blocks = nn.ModuleList(
        [
            _BasicTransformerBlock1D(
                dim=self.inner_dim,
                heads=config.num_attention_heads,
                dim_head=config.attention_head_dim,
                rope_type=config.rope_type,
                apply_gated_attention=config.apply_gated_attention,
            )
            for _ in range(config.num_layers)
        ]
    )
    self.num_learnable_registers = config.num_learnable_registers
    if self.num_learnable_registers:
        self.learnable_registers = nn.Parameter(
            torch.rand(
                self.num_learnable_registers,
                self.inner_dim,
                dtype=torch.bfloat16,
            )
            * 2.0
            - 1.0
        )

fastvideo.models.encoders.gemma.GemmaFeaturesExtractorProjLinear

GemmaFeaturesExtractorProjLinear(in_features: int, out_features: int)

Bases: Module

Linear projection that aggregates stacked Gemma hidden states.

Source code in fastvideo/models/encoders/gemma.py
def __init__(self, in_features: int, out_features: int) -> None:
    super().__init__()
    self.aggregate_embed = nn.Linear(in_features, out_features, bias=False)

fastvideo.models.encoders.gemma.LTX2GemmaTextEncoderModel

LTX2GemmaTextEncoderModel(config: TextEncoderConfig)

Bases: TextEncoder

Source code in fastvideo/models/encoders/gemma.py
def __init__(self, config: TextEncoderConfig) -> None:
    super().__init__(config)
    arch = config.arch_config

    # LTX-2.3 routes the caption projection before the connector and uses
    # separate per-modality feature extractor linears. Default (False)
    # keeps the single LTX-2.0 GemmaFeaturesExtractorProjLinear.
    self.use_v2_feature_extractor = bool(
        getattr(arch, "caption_proj_before_connector", False))
    if self.use_v2_feature_extractor:
        video_out_features = (
            getattr(arch, "video_feature_extractor_out_features", None)
            or arch.feature_extractor_out_features)
        audio_out_features = (
            getattr(arch, "audio_feature_extractor_out_features", None)
            or video_out_features)
        self.video_feature_extractor_linear = nn.Linear(
            arch.feature_extractor_in_features,
            video_out_features,
            bias=True,
        )
        self.audio_feature_extractor_linear = nn.Linear(
            arch.feature_extractor_in_features,
            audio_out_features,
            bias=True,
        )
    else:
        self.feature_extractor_linear = GemmaFeaturesExtractorProjLinear(
            in_features=arch.feature_extractor_in_features,
            out_features=arch.feature_extractor_out_features,
        )

    connector_apply_gated_attention = bool(
        getattr(arch, "connector_apply_gated_attention", False))
    video_connector_config = GemmaConnectorConfig(
        num_attention_heads=arch.connector_num_attention_heads,
        attention_head_dim=arch.connector_attention_head_dim,
        num_layers=arch.connector_num_layers,
        positional_embedding_theta=arch.connector_positional_embedding_theta,
        positional_embedding_max_pos=arch.connector_positional_embedding_max_pos,
        rope_type=LTXRopeType(arch.connector_rope_type),
        double_precision_rope=arch.connector_double_precision_rope,
        num_learnable_registers=arch.connector_num_learnable_registers,
        apply_gated_attention=connector_apply_gated_attention,
    )
    audio_connector_config = GemmaConnectorConfig(
        num_attention_heads=getattr(
            arch, "audio_connector_num_attention_heads", None)
        or arch.connector_num_attention_heads,
        attention_head_dim=getattr(
            arch, "audio_connector_attention_head_dim", None)
        or arch.connector_attention_head_dim,
        num_layers=getattr(arch, "audio_connector_num_layers", None)
        or arch.connector_num_layers,
        positional_embedding_theta=arch.connector_positional_embedding_theta,
        positional_embedding_max_pos=arch.connector_positional_embedding_max_pos,
        rope_type=LTXRopeType(arch.connector_rope_type),
        double_precision_rope=arch.connector_double_precision_rope,
        num_learnable_registers=arch.connector_num_learnable_registers,
        apply_gated_attention=connector_apply_gated_attention,
    )
    self.embeddings_connector = Embeddings1DConnector(video_connector_config)
    self.audio_embeddings_connector = Embeddings1DConnector(audio_connector_config)

    self.gemma_model_path = arch.gemma_model_path
    self.gemma_dtype = arch.gemma_dtype
    self.padding_side = arch.padding_side
    self._gemma_model: Gemma3ForConditionalGeneration | None = None

Methods:

fastvideo.models.encoders.gemma.LTX2GemmaTextEncoderModel.preprocess_text_embeddings
preprocess_text_embeddings(prompts: str | list[str], tokenizer: AutoTokenizer, tokenizer_kwargs: dict[str, Any] | None = None, padding_side: str | None = None) -> tuple[Tensor, Tensor]

Compute pre-connector text embeddings for LTX-2 training preprocessing.

Source code in fastvideo/models/encoders/gemma.py
@torch.no_grad()
def preprocess_text_embeddings(
    self,
    prompts: str | list[str],
    tokenizer: AutoTokenizer,
    tokenizer_kwargs: dict[str, Any] | None = None,
    padding_side: str | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Compute pre-connector text embeddings for LTX-2 training preprocessing."""
    if isinstance(prompts, str):
        prompts = [prompts]

    model = self.gemma_model
    kwargs: dict[str, Any] = {
        "padding": "max_length",
        "truncation": True,
        "return_tensors": "pt",
    }
    if tokenizer_kwargs is not None:
        kwargs.update(tokenizer_kwargs)
    if "max_length" not in kwargs:
        kwargs["max_length"] = self.config.arch_config.text_len

    original_padding_side = tokenizer.padding_side
    target_padding_side = padding_side or self.padding_side
    tokenizer.padding_side = target_padding_side
    try:
        text_inputs = tokenizer(prompts, **kwargs)
    finally:
        tokenizer.padding_side = original_padding_side

    input_ids = text_inputs["input_ids"].to(device=model.device)
    attention_mask = text_inputs["attention_mask"].to(device=model.device)
    outputs = model(
        input_ids=input_ids,
        attention_mask=attention_mask,
        output_hidden_states=True,
        return_dict=True,
    )
    video_features, _ = self._run_feature_extractor(
        outputs.hidden_states,
        attention_mask,
        padding_side=target_padding_side,
    )
    return video_features, attention_mask
fastvideo.models.encoders.gemma.LTX2GemmaTextEncoderModel.run_connectors
run_connectors(encoded_input: Tensor, attention_mask: Tensor) -> tuple[Tensor, Tensor, Tensor]

Apply embedding connectors to precomputed Gemma features.

Source code in fastvideo/models/encoders/gemma.py
def run_connectors(
    self,
    encoded_input: torch.Tensor,
    attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """Apply embedding connectors to precomputed Gemma features."""
    return self._run_connectors(encoded_input, encoded_input, attention_mask)

Functions: