"""Configuration for the AudioIn, Thinker, Talker and speech decoder.""" from __future__ import annotations from transformers import PretrainedConfig from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5Config from transformers.models.qwen3_omni_moe.configuration_qwen3_omni_moe import ( Qwen3OmniMoeAudioEncoderConfig, ) class EdgeInstantConfig(PretrainedConfig): model_type = "edgeinstant" sub_configs = { "thinker_config": Qwen3_5Config, "audio_config": Qwen3OmniMoeAudioEncoderConfig, } keys_to_ignore_at_inference = ["past_key_values"] def __init__( self, thinker_config=None, audio_config=None, projector_config=None, talker_config=None, codec_config=None, audio_tap="proj2", audio_pool=1, audio_stack=1, audio_special_ids=None, audio_pad_id=248076, sampling_rate=16000, output_sampling_rate=24000, max_audio_seconds=None, qa_prompt="Listen to the user's speech and answer.", thinker_special_ids=None, think_start_id=248068, think_end_id=248069, **kwargs, ): self.thinker_config = ( thinker_config if isinstance(thinker_config, Qwen3_5Config) else Qwen3_5Config(**(thinker_config or {})) ) self.audio_config = ( audio_config if isinstance(audio_config, Qwen3OmniMoeAudioEncoderConfig) else Qwen3OmniMoeAudioEncoderConfig(**(audio_config or {})) ) self.audio_tap = audio_tap self.audio_pool = int(audio_pool) self.audio_stack = int(audio_stack) if audio_tap not in {"proj2", "ln_post"}: raise ValueError(f"Unsupported AudioIn tower output: {audio_tap}") if self.audio_pool < 1 or self.audio_stack < 1: raise ValueError("audio_pool and audio_stack must be positive integers") audio_size = self.audio_config.d_model if audio_tap == "ln_post" else self.audio_config.output_dim hidden_size = self.thinker_config.text_config.hidden_size self.projector_config = { "norm_size": audio_size, "input_size": audio_size * self.audio_stack, "intermediate_size": None, "output_size": hidden_size, "align_mode": "replace", "audio_out_norm": False, **(projector_config or {}), } self.talker_config = dict(talker_config or {}) self.codec_config = dict(codec_config or {}) self.audio_special_ids = list(audio_special_ids if audio_special_ids is not None else [248070, 248071, 248076]) self.audio_pad_id = int(audio_pad_id) self.sampling_rate = int(sampling_rate) self.output_sampling_rate = int(output_sampling_rate) self.max_audio_seconds = max_audio_seconds self.qa_prompt = qa_prompt self.thinker_special_ids = list(thinker_special_ids or []) self.think_start_id = int(think_start_id) self.think_end_id = int(think_end_id) kwargs.setdefault("tie_word_embeddings", self.thinker_config.tie_word_embeddings) kwargs.setdefault("pad_token_id", self.thinker_config.text_config.pad_token_id) kwargs.setdefault("bos_token_id", self.thinker_config.text_config.bos_token_id) kwargs.setdefault("eos_token_id", self.thinker_config.text_config.eos_token_id) kwargs.setdefault("architectures", ["EdgeInstantForConditionalGeneration"]) kwargs.setdefault("processor_class", "EdgeInstantProcessor") kwargs.setdefault("auto_map", { "AutoConfig": "configuration_edgeinstant.EdgeInstantConfig", "AutoModel": "modeling_edgeinstant.EdgeInstantForConditionalGeneration", "AutoModelForCausalLM": "modeling_edgeinstant.EdgeInstantForConditionalGeneration", "AutoProcessor": "processing_edgeinstant.EdgeInstantProcessor", }) super().__init__(**kwargs) def get_text_config(self, decoder=None, encoder=None): return self.thinker_config.get_text_config(decoder=decoder, encoder=encoder) EdgeInstantConfig.register_for_auto_class()