CoolFace
Modelpublic

Firworks/Step-Audio-R1-nvfp4

sourceHugging Faceapache-2.0updated 10mo agoView on Hugging Face
2likes17downloads
configuration_step_audio_2.py129 linesDownload Raw Back to root
1from typing import Optional, Union2 3from transformers import Qwen2Config4from transformers.configuration_utils import PretrainedConfig5 6 7class StepAudio2EncoderConfig(PretrainedConfig):8    model_type = "step_audio_2_encoder"9 10    def __init__(11        self,12        n_mels=128,13        n_audio_ctx=1500,14        n_audio_state=512,15        n_audio_head=8,16        n_audio_layer=6,17        llm_dim=4096,18        kernel_size=3,19        adapter_stride=2,20        **kwargs,21    ):22        self.n_mels      = n_mels23        self.n_audio_ctx = n_audio_ctx24        self.n_audio_state = n_audio_state25        self.n_audio_head = n_audio_head26        self.n_audio_layer = n_audio_layer27        self.llm_dim     = llm_dim28        self.kernel_size = kernel_size29        self.adapter_stride = adapter_stride30        super().__init__(**kwargs)31 32class StepAudio2TextConfig(PretrainedConfig):33    model_type = "step_audio_2_text"34 35    def __init__(36        self,37        vocab_size=64012,38        hidden_size=4096,39        intermediate_size=11008,40        num_hidden_layers=48,41        num_attention_heads=32,42        num_attention_groups=4,43        num_key_value_heads=4,44        hidden_act="silu",45        max_position_embeddings=8192,46        initializer_range=0.02,47        rms_norm_eps=1e-6,48        rope_theta=1000000.0,49        rope_scaling=None,50        eos_token_id=None,51        **kwargs52    ):53 54        if eos_token_id is not None:55            if isinstance(eos_token_id, list):56                eos_token_id = list(set([151643, 151645, 151665] + eos_token_id))57            else:58                eos_token_id = [151643, 151645, 151665, eos_token_id]59        else:60            eos_token_id = [151643, 151645, 151665]61 62        super().__init__(63            eos_token_id=eos_token_id,64            **kwargs)65 66        self.vocab_size = vocab_size67        self.hidden_size = hidden_size68        self.intermediate_size = intermediate_size69        self.num_hidden_layers = num_hidden_layers70        self.num_attention_heads = num_attention_heads71        self.num_attention_groups = num_attention_groups72        self.num_key_value_heads = num_key_value_heads73        assert self.num_attention_groups == self.num_key_value_heads, "num_attention_groups must be equal to num_key_value_heads"74        self.hidden_act = hidden_act75        self.max_position_embeddings = max_position_embeddings76        self.initializer_range = initializer_range77        self.rms_norm_eps = rms_norm_eps78        self.rope_theta = rope_theta79        self.rope_scaling = rope_scaling80 81        self.text_config = Qwen2Config(82            vocab_size=vocab_size,83            hidden_size=hidden_size,84            intermediate_size=intermediate_size,85            num_hidden_layers=num_hidden_layers,86            num_attention_heads=num_attention_heads,87            num_key_value_heads=num_key_value_heads,88            hidden_act=hidden_act,89            max_position_embeddings=max_position_embeddings,90            initializer_range=initializer_range,91            rms_norm_eps=rms_norm_eps,92            rope_theta=rope_theta,93            rope_scaling=rope_scaling,94            architectures=["Qwen2ForCausalLM"],95            torch_dtype=getattr(self, "torch_dtype", "bfloat16"),96        )97 98class StepAudio2Config(PretrainedConfig):99    model_type = "step_audio_2"100    architectures = ["StepAudio2ForCausalLM"]101 102    def __init__(103        self,104        audio_encoder_config :Optional[Union[dict, StepAudio2EncoderConfig]] = None,105        text_config: Optional[Union[dict, StepAudio2TextConfig]] = None,106        use_sliding_window: bool = False,107        sliding_window: Optional[int] = 2048,108        max_window_layers: Optional[int] = None,109        **kwargs110    ):111        kwargs.setdefault("use_sliding_window", use_sliding_window)112        kwargs.setdefault("sliding_window", sliding_window)113        if max_window_layers is None:114            max_window_layers = kwargs.get("num_hidden_layers", None)115        kwargs.setdefault("max_window_layers", max_window_layers)116        super().__init__(**kwargs)117 118        if text_config is None:119            text_config = StepAudio2TextConfig().text_config120        elif isinstance(text_config, dict):121            text_config = StepAudio2TextConfig(**text_config).text_config122 123        self.text_config = text_config124 125        if audio_encoder_config is None:126            self.audio_encoder_config = StepAudio2EncoderConfig()127        elif isinstance(audio_encoder_config, dict):128            self.audio_encoder_config = StepAudio2EncoderConfig(**audio_encoder_config)129