CoolFace
Modelpublic

TransWithAI/Step-Audio-R1.1-NVFP4

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
1likes12downloads
configuration_step_audio_2.py176 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        # Get torch_dtype from kwargs if provided82        torch_dtype = kwargs.get("torch_dtype", getattr(self, "torch_dtype", "bfloat16"))83        84        self.text_config = Qwen2Config(85            vocab_size=vocab_size,86            hidden_size=hidden_size,87            intermediate_size=intermediate_size,88            num_hidden_layers=num_hidden_layers,89            num_attention_heads=num_attention_heads,90            num_key_value_heads=num_key_value_heads,91            hidden_act=hidden_act,92            max_position_embeddings=max_position_embeddings,93            initializer_range=initializer_range,94            rms_norm_eps=rms_norm_eps,95            rope_theta=rope_theta,96            rope_scaling=rope_scaling,97            architectures=["Qwen2ForCausalLM"],98            torch_dtype=torch_dtype,99        )100 101class StepAudio2Config(PretrainedConfig):102    model_type = "step_audio_2"103    architectures = ["StepAudio2ForCausalLM"]104    105    # Support alternative model types and architectures for 32B model106    # This allows the config to work with both "step_audio_2" and "step_audio_qwen2" model types107 108    def __init__(109        self,110        audio_encoder_config :Optional[Union[dict, StepAudio2EncoderConfig]] = None,111        text_config: Optional[Union[dict, StepAudio2TextConfig]] = None,112        use_sliding_window: bool = False,113        sliding_window: Optional[int] = 2048,114        max_window_layers: Optional[int] = None,115        **kwargs116    ):117        kwargs.setdefault("use_sliding_window", use_sliding_window)118        kwargs.setdefault("sliding_window", sliding_window)119        if max_window_layers is None:120            max_window_layers = kwargs.get("num_hidden_layers", None)121        kwargs.setdefault("max_window_layers", max_window_layers)122        123        # Save torch_dtype if provided (for 32B model flat config)124        if 'torch_dtype' in kwargs:125            self.torch_dtype = kwargs['torch_dtype']126        127        super().__init__(**kwargs)128 129        # Support for flat config structure (32B model format)130        # If text_config is None and we have flat config parameters, extract them131        if text_config is None:132            # Check if we have flat config parameters (32B model format)133            flat_text_params = {}134            text_param_names = [135                'vocab_size', 'hidden_size', 'intermediate_size', 'num_hidden_layers',136                'num_attention_heads', 'num_attention_groups', 'num_key_value_heads',137                'hidden_act', 'max_position_embeddings', 'initializer_range',138                'rms_norm_eps', 'rope_theta', 'rope_scaling', 'eos_token_id', 'pad_token_id'139            ]140            141            for param_name in text_param_names:142                if param_name in kwargs:143                    flat_text_params[param_name] = kwargs[param_name]144            145            # Set default hidden_act if not present (32B model config doesn't have it)146            if 'hidden_act' not in flat_text_params:147                flat_text_params['hidden_act'] = 'silu'148            149            # Set default initializer_range if not present150            if 'initializer_range' not in flat_text_params:151                flat_text_params['initializer_range'] = 0.02152            153            # Also check for torch_dtype154            if 'torch_dtype' in kwargs:155                flat_text_params['torch_dtype'] = kwargs['torch_dtype']156            157            if flat_text_params:158                # We have flat config, use it to build text_config159                text_config = StepAudio2TextConfig(**flat_text_params).text_config160            else:161                # No flat config, use defaults162                text_config = StepAudio2TextConfig().text_config163        elif isinstance(text_config, dict):164            text_config = StepAudio2TextConfig(**text_config).text_config165 166        self.text_config = text_config167 168        if audio_encoder_config is None:169            # Check if we have flat audio_encoder_config in kwargs170            if 'audio_encoder_config' in kwargs and isinstance(kwargs['audio_encoder_config'], dict):171                self.audio_encoder_config = StepAudio2EncoderConfig(**kwargs['audio_encoder_config'])172            else:173                self.audio_encoder_config = StepAudio2EncoderConfig()174        elif isinstance(audio_encoder_config, dict):175            self.audio_encoder_config = StepAudio2EncoderConfig(**audio_encoder_config)176