""" SoriSpeech Model for Conditional Generation (Speech-to-Text) This model uses Qwen3OmniMoe-compatible AudioEncoder for audio processing and Qwen3 as the backbone LLM for text generation. The Audio Encoder architecture is designed to be weight-compatible with Qwen3-Omni-30B-A3B-Instruct's audio_tower. """ import math from dataclasses import dataclass from typing import Optional, Tuple, Union, List import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from torch.nn import CrossEntropyLoss from transformers import ( PreTrainedModel, PretrainedConfig, AutoModelForCausalLM, AutoConfig, GenerationMixin, ) from transformers.activations import ACT2FN from transformers.modeling_outputs import ( BaseModelOutput, CausalLMOutputWithPast, ) from transformers.utils import logging logger = logging.get_logger(__name__) # ============================================================================ # Audio Encoder Config (compatible with Qwen3OmniMoeAudioEncoder) # ============================================================================ class SoriAudioEncoderConfig(PretrainedConfig): """ Configuration for the audio encoder. Compatible with Qwen3OmniMoeAudioEncoderConfig. """ model_type = "sori_audio_encoder" def __init__( self, d_model: int = 1280, encoder_layers: int = 32, encoder_attention_heads: int = 20, encoder_ffn_dim: int = 5120, dropout: float = 0.0, attention_dropout: float = 0.0, activation_function: str = "gelu", num_mel_bins: int = 128, max_source_positions: int = 1500, scale_embedding: bool = False, n_window: int = 50, n_window_infer: int = 800, conv_chunksize: int = 500, downsample_hidden_size: int = 480, output_dim: int = 2048, **kwargs, ): super().__init__(**kwargs) self.d_model = d_model self.encoder_layers = encoder_layers self.encoder_attention_heads = encoder_attention_heads self.encoder_ffn_dim = encoder_ffn_dim self.dropout = dropout self.attention_dropout = attention_dropout self.activation_function = activation_function self.num_mel_bins = num_mel_bins self.max_source_positions = max_source_positions self.scale_embedding = scale_embedding self.n_window = n_window self.n_window_infer = n_window_infer self.conv_chunksize = conv_chunksize self.downsample_hidden_size = downsample_hidden_size self.output_dim = output_dim class SoriSpeechConfig(PretrainedConfig): """Configuration class for SoriSpeech model.""" model_type = "sori_speech" is_composition = True def __init__( self, audio_encoder_config: Optional[dict] = None, text_config: Optional[dict] = None, llm_model_name_or_path: str = "Qwen/Qwen3-4B-Instruct-2507", audio_token_id: int = 151669, audio_start_token_id: int = 151670, audio_end_token_id: int = 151671, pad_token_id: int = 151643, bos_token_id: int = 151643, eos_token_id: int = 151645, vocab_size: int = 151672, audio_to_llm_proj: bool = True, **kwargs, ): super().__init__( pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, **kwargs, ) if audio_encoder_config is None: audio_encoder_config = {} if isinstance(audio_encoder_config, dict): self.audio_encoder_config = SoriAudioEncoderConfig(**audio_encoder_config) else: self.audio_encoder_config = audio_encoder_config self._text_config_dict = text_config self.llm_model_name_or_path = llm_model_name_or_path self.audio_token_id = audio_token_id self.audio_start_token_id = audio_start_token_id self.audio_end_token_id = audio_end_token_id self.vocab_size = vocab_size self.audio_to_llm_proj = audio_to_llm_proj @property def text_config(self): return self._text_config_dict @text_config.setter def text_config(self, value): self._text_config_dict = value def get_text_config(self, decoder=False): if self._text_config_dict is not None: from transformers import Qwen3Config return Qwen3Config(**self._text_config_dict) return self def get_decoder(self): return self.get_text_config(decoder=True) # ============================================================================ # Audio Encoder Components (Qwen3OmniMoe compatible) # ============================================================================ def _get_feat_extract_output_lengths(input_lengths: torch.Tensor) -> torch.Tensor: """Calculate output lengths after conv layers.""" input_lengths_leave = input_lengths % 100 feat_lengths = (input_lengths_leave - 1) // 2 + 1 output_lengths = ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (input_lengths // 100) * 13 return output_lengths class SinusoidsPositionEmbedding(nn.Module): """Sinusoidal position embedding (same as Qwen3OmniMoe).""" def __init__(self, length: int, channels: int, max_timescale: int = 10000): super().__init__() self.length = length self.channels = channels self.max_timescale = max_timescale if channels % 2 != 0: raise ValueError("SinusoidsPositionEmbedding needs even channels input") log_timescale_increment = np.log(max_timescale) / (channels // 2 - 1) inv_timescales = torch.exp(-log_timescale_increment * torch.arange(channels // 2).float()) scaled_time = torch.arange(length)[:, np.newaxis] * inv_timescales[np.newaxis, :] self.register_buffer( "positional_embedding", torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], dim=1), persistent=False, ) def forward(self, seqlen: int) -> torch.Tensor: return self.positional_embedding[:seqlen, :] class SoriAudioAttention(nn.Module): """Audio attention layer (compatible with Qwen3OmniMoeAudioAttention).""" def __init__(self, config: SoriAudioEncoderConfig): super().__init__() self.embed_dim = config.d_model self.num_heads = config.encoder_attention_heads self.dropout = config.attention_dropout self.head_dim = self.embed_dim // self.num_heads self.config = config self.scaling = self.head_dim ** -0.5 self.k_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) self.v_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) self.q_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) self.out_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) def forward(self, hidden_states, cu_seqlens=None, attention_mask=None, **kwargs): seq_length, _ = hidden_states.size() query_states = self.q_proj(hidden_states).reshape(seq_length, self.num_heads, -1) key_states = self.k_proj(hidden_states).reshape(seq_length, self.num_heads, -1) value_states = self.v_proj(hidden_states).reshape(seq_length, self.num_heads, -1) query_states = query_states.transpose(0, 1).unsqueeze(0) key_states = key_states.transpose(0, 1).unsqueeze(0) value_states = value_states.transpose(0, 1).unsqueeze(0) attn_weights = torch.matmul(query_states, key_states.transpose(-1, -2)) * self.scaling if attention_mask is not None: attn_weights = attn_weights + attention_mask attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype) attn_weights = F.dropout(attn_weights, p=self.dropout if self.training else 0.0, training=self.training) attn_output = torch.matmul(attn_weights, value_states) attn_output = attn_output.squeeze(0).transpose(0, 1).reshape(seq_length, -1).contiguous() attn_output = self.out_proj(attn_output) return attn_output class SoriAudioEncoderLayer(nn.Module): """Audio encoder layer (compatible with Qwen3OmniMoeAudioEncoderLayer).""" def __init__(self, config: SoriAudioEncoderConfig): super().__init__() self.embed_dim = config.d_model self.self_attn = SoriAudioAttention(config) self.self_attn_layer_norm = nn.LayerNorm(self.embed_dim) self.dropout = config.dropout self.activation_fn = ACT2FN[config.activation_function] self.activation_dropout = config.dropout self.fc1 = nn.Linear(self.embed_dim, config.encoder_ffn_dim) self.fc2 = nn.Linear(config.encoder_ffn_dim, self.embed_dim) self.final_layer_norm = nn.LayerNorm(self.embed_dim) def forward(self, hidden_states, cu_seqlens, attention_mask=None, **kwargs): residual = hidden_states hidden_states = self.self_attn_layer_norm(hidden_states) hidden_states = self.self_attn(hidden_states, cu_seqlens, attention_mask, **kwargs) hidden_states = F.dropout(hidden_states, p=self.dropout, training=self.training) hidden_states = residual + hidden_states residual = hidden_states hidden_states = self.final_layer_norm(hidden_states) hidden_states = self.fc1(hidden_states) hidden_states = self.activation_fn(hidden_states) hidden_states = F.dropout(hidden_states, p=self.activation_dropout, training=self.training) hidden_states = self.fc2(hidden_states) hidden_states = F.dropout(hidden_states, p=self.dropout, training=self.training) hidden_states = residual + hidden_states return (hidden_states,) class SoriAudioEncoder(nn.Module): """ Audio Encoder compatible with Qwen3OmniMoeAudioEncoder. """ def __init__(self, config: SoriAudioEncoderConfig): super().__init__() self.config = config self.dropout = config.dropout embed_dim = config.d_model self.num_mel_bins = config.num_mel_bins self.max_source_positions = config.max_source_positions self.embed_scale = math.sqrt(embed_dim) if config.scale_embedding else 1.0 self.n_window = config.n_window self.positional_embedding = SinusoidsPositionEmbedding(self.max_source_positions, embed_dim) self.layers = nn.ModuleList([ SoriAudioEncoderLayer(config) for _ in range(config.encoder_layers) ]) self.ln_post = nn.LayerNorm(config.d_model) self.conv2d1 = nn.Conv2d(1, config.downsample_hidden_size, 3, 2, padding=1) self.conv2d2 = nn.Conv2d(config.downsample_hidden_size, config.downsample_hidden_size, 3, 2, padding=1) self.conv2d3 = nn.Conv2d(config.downsample_hidden_size, config.downsample_hidden_size, 3, 2, padding=1) mel_after_conv = (((config.num_mel_bins + 1) // 2 + 1) // 2 + 1) // 2 self.conv_out = nn.Linear(config.downsample_hidden_size * mel_after_conv, config.d_model, bias=False) self.proj1 = nn.Linear(config.d_model, config.d_model) self.act = ACT2FN[config.activation_function] self.proj2 = nn.Linear(config.d_model, config.output_dim) self.n_window_infer = config.n_window_infer self.conv_chunksize = config.conv_chunksize def _prepare_attention_mask(self, inputs_tensor, cu_seqlens): seq_length = inputs_tensor.shape[0] attention_mask = torch.full( [1, 1, seq_length, seq_length], torch.finfo(inputs_tensor.dtype).min, device=inputs_tensor.device, dtype=inputs_tensor.dtype, ) for i in range(1, len(cu_seqlens)): attention_mask[..., cu_seqlens[i - 1]:cu_seqlens[i], cu_seqlens[i - 1]:cu_seqlens[i]] = 0 return attention_mask def forward(self, input_features, feature_lens, **kwargs): aftercnn_lens = _get_feat_extract_output_lengths(feature_lens) chunk_num = torch.ceil(feature_lens / (self.n_window * 2)).long() chunk_lengths = torch.tensor( [self.n_window * 2] * chunk_num.sum(), dtype=torch.long, device=feature_lens.device ) tail_chunk_index = F.pad(chunk_num, (1, 0), value=-1).cumsum(0)[1:] chunk_lengths[tail_chunk_index] = feature_lens % (self.n_window * 2) chunk_lengths[chunk_lengths == 0] = self.n_window * 2 chunk_list = input_features.T.split(chunk_lengths.tolist(), dim=0) padded_feature = nn.utils.rnn.pad_sequence(chunk_list, batch_first=True).transpose(1, 2) feature_lens_after_cnn = _get_feat_extract_output_lengths(chunk_lengths) padded_mask_after_cnn = nn.utils.rnn.pad_sequence( [torch.ones(length, dtype=torch.bool, device=padded_feature.device) for length in feature_lens_after_cnn], batch_first=True, ) padded_feature = padded_feature.unsqueeze(1) padded_embeds = [] for chunk in padded_feature.split(self.conv_chunksize, dim=0): padded_embed = F.gelu(self.conv2d1(chunk)) padded_embed = F.gelu(self.conv2d2(padded_embed)) padded_embed = F.gelu(self.conv2d3(padded_embed)) padded_embeds.append(padded_embed) padded_embed = torch.cat(padded_embeds, dim=0) b, c, f, t = padded_embed.size() padded_embed = self.conv_out(padded_embed.permute(0, 3, 1, 2).contiguous().view(b, t, c * f)) positional_embedding = self.positional_embedding.positional_embedding[:padded_embed.shape[1], :].unsqueeze(0).to(padded_embed.dtype) padded_embed = padded_embed + positional_embedding hidden_states = padded_embed[padded_mask_after_cnn] cu_chunk_lens = [0] window_aftercnn = padded_mask_after_cnn.shape[-1] * (self.n_window_infer // (self.n_window * 2)) for cnn_len in aftercnn_lens: cu_chunk_lens += [window_aftercnn] * (cnn_len // window_aftercnn) remainder = cnn_len % window_aftercnn if remainder != 0: cu_chunk_lens += [remainder] cu_seqlens = torch.tensor(cu_chunk_lens, device=aftercnn_lens.device).cumsum(-1, dtype=torch.int32) attention_mask = self._prepare_attention_mask(hidden_states, cu_seqlens) for encoder_layer in self.layers: layer_outputs = encoder_layer(hidden_states, cu_seqlens, attention_mask=attention_mask) hidden_states = layer_outputs[0] hidden_states = self.ln_post(hidden_states) hidden_states = self.proj1(hidden_states) hidden_states = self.act(hidden_states) hidden_states = self.proj2(hidden_states) return BaseModelOutput(last_hidden_state=hidden_states), aftercnn_lens # ============================================================================ # SoriSpeech Model # ============================================================================ class SoriSpeechPreTrainedModel(PreTrainedModel): config_class = SoriSpeechConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["SoriAudioEncoderLayer"] def _init_weights(self, module): std = 0.02 if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=std) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=std) elif isinstance(module, nn.Conv2d): nn.init.kaiming_normal_(module.weight, mode="fan_out", nonlinearity="relu") if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.LayerNorm): module.bias.data.zero_() module.weight.data.fill_(1.0) class SoriSpeechForConditionalGeneration(SoriSpeechPreTrainedModel, GenerationMixin): """ SoriSpeech model for speech-to-text generation. Combines: - Qwen3OmniMoe-compatible Audio Encoder - Qwen3 Language Model as decoder """ _tied_weights_keys = ["language_model.lm_head.weight"] def __init__(self, config: SoriSpeechConfig): super().__init__(config) self.config = config # Audio encoder (Qwen3OmniMoe compatible) self.audio_encoder = SoriAudioEncoder(config.audio_encoder_config) # Language model - initialize from text_config (NOT from pretrained) self.language_model = None self._init_language_model_from_config(config) # Optional projection from audio features to LLM hidden size self.audio_proj = None if config.audio_to_llm_proj and self.language_model is not None: llm_hidden_size = self.language_model.config.hidden_size audio_output_dim = config.audio_encoder_config.output_dim if llm_hidden_size != audio_output_dim: self.audio_proj = nn.Linear(audio_output_dim, llm_hidden_size, bias=False) def _init_language_model_from_config(self, config): """ Initialize language model from text_config. This creates an empty model with correct architecture. Weights will be loaded by from_pretrained(). """ if config.text_config is not None: try: from transformers import Qwen3ForCausalLM, Qwen3Config text_config_dict = config.text_config.copy() # Use vocab_size from main config if config.vocab_size is not None: text_config_dict['vocab_size'] = config.vocab_size text_config = Qwen3Config(**text_config_dict) self.language_model = Qwen3ForCausalLM(text_config) logger.info(f"Initialized language model from text_config (vocab_size={text_config.vocab_size})") except Exception as e: logger.warning(f"Could not initialize from text_config: {e}") raise else: # If no text_config, we need to load from pretrained # This is only for initial model creation, not for from_pretrained loading self._load_pretrained_llm(config) def _load_pretrained_llm(self, config): """Load pretrained LLM (only for initial model creation).""" try: self.language_model = AutoModelForCausalLM.from_pretrained( config.llm_model_name_or_path, torch_dtype=torch.bfloat16, trust_remote_code=True, ) # Save the config for future use config._text_config_dict = self.language_model.config.to_dict() logger.info(f"Loaded LLM from {config.llm_model_name_or_path}") except Exception as e: raise RuntimeError(f"Failed to load language model: {e}") def get_input_embeddings(self): return self.language_model.get_input_embeddings() def set_input_embeddings(self, value): self.language_model.set_input_embeddings(value) def get_output_embeddings(self): return self.language_model.get_output_embeddings() def set_output_embeddings(self, new_embeddings): self.language_model.set_output_embeddings(new_embeddings) def resize_token_embeddings(self, new_num_tokens): return self.language_model.resize_token_embeddings(new_num_tokens) def forward( self, input_ids=None, input_features=None, feature_lens=None, attention_mask=None, position_ids=None, past_key_values=None, inputs_embeds=None, labels=None, use_cache=None, output_attentions=None, output_hidden_states=None, return_dict=None, **kwargs, ): # If no audio, just forward to language model if input_features is None: return self.language_model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, labels=labels, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, **kwargs, ) # Encode audio audio_output, audio_feature_lens = self.audio_encoder( input_features=input_features, feature_lens=feature_lens, ) audio_features = audio_output.last_hidden_state # Project audio features to LLM hidden size if needed if self.audio_proj is not None: audio_features = self.audio_proj(audio_features) # Get text embeddings if inputs_embeds is None and input_ids is not None: inputs_embeds = self.get_input_embeddings()(input_ids) # Merge audio and text if input_ids is not None: inputs_embeds = self._merge_audio_and_text( input_ids, inputs_embeds, audio_features, audio_feature_lens ) return self.language_model( input_ids=None, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, labels=labels, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, **kwargs, ) def _merge_audio_and_text(self, input_ids, inputs_embeds, audio_features, audio_feature_lens): """Merge audio features into text embeddings at audio token positions.""" batch_size = input_ids.shape[0] audio_token_mask = input_ids == self.config.audio_token_id merged_embeds = inputs_embeds.clone() audio_offset = 0 for b in range(batch_size): audio_positions = torch.where(audio_token_mask[b])[0] if len(audio_positions) > 0: audio_len = int(audio_feature_lens[b].item()) if audio_feature_lens.dim() > 0 else int(audio_feature_lens.item()) batch_audio = audio_features[audio_offset:audio_offset + audio_len] audio_offset += audio_len num_tokens = len(audio_positions) if num_tokens == audio_len: merged_embeds[b, audio_positions] = batch_audio.to(merged_embeds.dtype) elif audio_len > num_tokens: merged_embeds[b, audio_positions] = batch_audio[:num_tokens].to(merged_embeds.dtype) else: merged_embeds[b, audio_positions[:audio_len]] = batch_audio.to(merged_embeds.dtype) return merged_embeds def generate(self, **kwargs): return self.language_model.generate(**kwargs) def prepare_inputs_for_generation(self, *args, **kwargs): return self.language_model.prepare_inputs_for_generation(*args, **kwargs) @staticmethod def _reorder_cache(past_key_values, beam_idx): reordered_past = () for layer_past in past_key_values: reordered_past += (tuple(p.index_select(0, beam_idx.to(p.device)) for p in layer_past),) return reordered_past # Register for auto classes AutoConfig.register("sori_speech", SoriSpeechConfig) if __name__ == "__main__": print("Testing SoriSpeech components...") audio_config = SoriAudioEncoderConfig( d_model=1280, encoder_layers=4, encoder_attention_heads=20, encoder_ffn_dim=5120, num_mel_bins=128, output_dim=2048, downsample_hidden_size=480, n_window=50, n_window_infer=800, conv_chunksize=500, ) audio_encoder = SoriAudioEncoder(audio_config) print(f"Audio encoder: {sum(p.numel() for p in audio_encoder.parameters()):,} params") dummy_mel = torch.randn(128, 500) dummy_lens = torch.tensor([500]) output, out_lens = audio_encoder(dummy_mel, dummy_lens) print(f"Input shape: (128, 500)") print(f"Output shape: {output.last_hidden_state.shape}") print(f"Output lens: {out_lens}") print("Done!")