gemma-4-E4B-it-audio-encoder / gemma4_audio_encoder.py
Aratako's picture
Update gemma4_audio_encoder.py
8f87fc0 verified
Raw
History Blame Contribute Delete
1.4 kB
from transformers import PreTrainedModel
from transformers.models.gemma4.configuration_gemma4 import Gemma4Config
from transformers.models.gemma4.modeling_gemma4 import Gemma4AudioModel, Gemma4MultimodalEmbedder
class Gemma4AudioEncoder(PreTrainedModel):
config_class = Gemma4Config
def __init__(self, config):
super().__init__(config)
self.audio_tower = Gemma4AudioModel(config.audio_config)
self.embed_audio = Gemma4MultimodalEmbedder(config.audio_config, config.text_config)
self.post_init()
def forward(self, input_features, input_features_mask, project=True, **kwargs):
"""
Args:
input_features: Audio mel-spectrogram features.
input_features_mask: Attention mask for audio features (True = valid, False = padding).
project: If True, project to LLM embedding space (2560-dim).
If False, return audio tower output (1536-dim).
Returns:
If project=True: (projected_features, attention_mask)
If project=False: (encoder_features, attention_mask)
"""
output = self.audio_tower(input_features, input_features_mask)
if project:
projected = self.embed_audio(inputs_embeds=output.last_hidden_state)
return projected, output.attention_mask
return output.last_hidden_state, output.attention_mask