TransWithAI/Step-Audio-R1.1-NVFP4
112
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 