Firworks/Step-Audio-R1-nvfp4
217
1from typing import Iterable, Optional, Tuple2 3import librosa4import torch5import torch.nn.functional as F6import torchaudio7from torch import Tensor, nn8from transformers import PreTrainedModel, Qwen2Model9from transformers.generation.utils import GenerationMixin10from transformers.modeling_outputs import CausalLMOutputWithPast11 12from .configuration_step_audio_2 import StepAudio2Config13 14 15def _mel_filters(n_mels: int) -> torch.Tensor:16 """Load the mel filterbank matrix for projecting STFT into a Mel spectrogram."""17 assert n_mels in {80, 128}, f"Unsupported n_mels: {n_mels}"18 if n_mels == 128:19 return torch.from_numpy(librosa.filters.mel(sr=16000, n_fft=400, n_mels=128))20 else:21 return torch.from_numpy(librosa.filters.mel(sr=16000, n_fft=400, n_mels=80))22 23 24def load_audio(file_path, target_rate=16000, max_length=None):25 """26 Open an audio file and read as mono waveform, resampling as necessary27 If max_length is provided, truncate the audio to that length28 """29 waveform, sample_rate = torchaudio.load(file_path)30 if sample_rate != target_rate:31 waveform = torchaudio.transforms.Resample(orig_freq=sample_rate, new_freq=target_rate)(waveform)32 audio = waveform[0] # get the first channel33 34 # Truncate audio if it exceeds max_length35 if max_length is not None and audio.shape[0] > max_length:36 audio = audio[:max_length]37 38 return audio39 40def log_mel_spectrogram(audio, n_mels=128, padding=479, device=None):41 """42 Compute the log-Mel spectrogram with specific padding for StepAudio43 """44 if not torch.is_tensor(audio):45 if isinstance(audio, str):46 audio = load_audio(audio)47 audio = torch.from_numpy(audio)48 if device is not None:49 audio = audio.to(device)50 if padding > 0:51 audio = F.pad(audio, (0, padding))52 window = torch.hann_window(400).to(audio.device)53 stft = torch.stft(audio, 400, 160, window=window, return_complex=True)54 magnitudes = stft[..., :-1].abs() ** 255 filters = _mel_filters(n_mels)56 mel_spec = filters @ magnitudes57 58 log_spec = torch.clamp(mel_spec, min=1e-10).log10()59 log_spec = torch.maximum(log_spec, log_spec.max() - 8.0)60 log_spec = (log_spec + 4.0) / 4.061 return log_spec62 63def compute_token_num(max_feature_len):64 # First, audio goes through encoder:65 # 1. conv1: kernel=3, stride=1, padding=1 -> size unchanged66 # 2. conv2: kernel=3, stride=2, padding=1 -> size/267 # 3. avg_pooler: kernel=2, stride=2 -> size/268 max_feature_len = max_feature_len - 2 # remove padding69 encoder_output_dim = (max_feature_len + 1) // 2 // 2 # after conv2 and avg_pooler70 71 # Then through adaptor (parameters from config file):72 padding = 173 kernel_size = 3 # from config: audio_encoder_config.kernel_size74 stride = 2 # from config: audio_encoder_config.adapter_stride75 adapter_output_dim = (encoder_output_dim + 2 * padding - kernel_size) // stride + 176 return adapter_output_dim77 78def make_non_pad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor:79 """Make mask tensor containing indices of non-padded part.80 81 The sequences in a batch may have different lengths. To enable82 batch computing, padding is need to make all sequence in same83 size. To avoid the padding part pass value to context dependent84 block such as attention or convolution , this padding part is85 masked.86 87 1 for non-padded part and 0 for padded part.88 89 Parameters90 ----------91 lengths (torch.Tensor): Batch of lengths (B,).92 93 Returns:94 -------95 torch.Tensor: Mask tensor containing indices of padded part (B, max_T).96 97 Examples:98 >>> import torch99 >>> import s3tokenizer100 >>> lengths = torch.tensor([5, 3, 2])101 >>> masks = s3tokenizer.make_non_pad_mask(lengths)102 masks = [[1, 1, 1, 1, 1],103 [1, 1, 1, 0, 0],104 [1, 1, 0, 0, 0]]105 """106 batch_size = lengths.size(0)107 max_len = max_len if max_len > 0 else lengths.max().item()108 seq_range = torch.arange(0,109 max_len,110 dtype=torch.int64,111 device=lengths.device)112 seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len)113 seq_length_expand = lengths.unsqueeze(-1)114 mask = seq_range_expand >= seq_length_expand115 return ~mask116 117def mask_to_bias(mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:118 """Convert bool-tensor to float-tensor for flash attention.119 120 Parameters121 ----------122 lengths (torch.Tensor): Batch of lengths (B, ?).123 124 Returns:125 -------126 torch.Tensor: Mask tensor containing indices of padded part (B, ?).127 128 Examples:129 >>> import torch130 >>> import s3tokenizer131 >>> lengths = torch.tensor([5, 3, 2])132 >>> masks = s3tokenizer.make_non_pad_mask(lengths)133 masks = [[1, 1, 1, 1, 1],134 [1, 1, 1, 0, 0],135 [1, 1, 0, 0, 0]]136 >>> new_masks = s3tokenizer.mask_to_bias(masks, torch.float32)137 new_masks = [[-0.0000e+00, -0.0000e+00, -0.0000e+00, -0.0000e+00, -0.0000e+00],138 [-0.0000e+00, -0.0000e+00, -0.0000e+00, -1.0000e+10, -1.0000e+10],139 [-0.0000e+00, -0.0000e+00, -1.0000e+10, -1.0000e+10, -1.0000e+10]]140 """141 assert mask.dtype == torch.bool142 assert dtype in [torch.float32, torch.bfloat16, torch.float16]143 mask = mask.to(dtype)144 # attention mask bias145 # NOTE(Mddct): torch.finfo jit issues146 # chunk_masks = (1.0 - chunk_masks) * torch.finfo(dtype).min147 mask = (1.0 - mask) * -1.0e+10148 return mask149 150class LayerNorm(nn.LayerNorm):151 def forward(self, input: Tensor) -> Tensor:152 return super().forward(input).type(input.dtype)153 154class Linear(nn.Linear):155 def forward(self, input: Tensor) -> Tensor:156 return F.linear(157 input,158 self.weight.to(input.dtype),159 None if self.bias is None else self.bias.to(input.dtype),160 )161 162class Conv1d(nn.Conv1d):163 def _conv_forward(164 self, input: Tensor, weight: Tensor, bias: Optional[Tensor]165 ) -> Tensor:166 return super()._conv_forward(167 input, weight.to(input.dtype), None if bias is None else bias.to(input.dtype)168 )169 170class MultiHeadAttention(nn.Module):171 def __init__(self, n_state: int, n_head: int):172 super().__init__()173 self.n_head = n_head174 self.query = Linear(n_state, n_state)175 self.key = Linear(n_state, n_state, bias=False)176 self.value = Linear(n_state, n_state)177 self.out = Linear(n_state, n_state)178 179 def forward(180 self,181 x: Tensor,182 mask: Optional[Tensor] = None,183 ):184 q = self.query(x)185 k = self.key(x)186 v = self.value(x)187 188 wv, qk = self.qkv_attention(q, k, v, mask)189 return self.out(wv), qk190 191 def qkv_attention(192 self, q: Tensor, k: Tensor, v: Tensor, mask: Optional[Tensor] = None193 ):194 _, T, D = q.shape195 scale = (D // self.n_head) ** -0.25196 q = q.view(*q.shape[:2], self.n_head, -1).permute(0, 2, 1, 3) * scale197 k = k.view(*k.shape[:2], self.n_head, -1).permute(0, 2, 3, 1) * scale198 v = v.view(*v.shape[:2], self.n_head, -1).permute(0, 2, 1, 3)199 200 qk = q @ k # (B, n_head, T, T)201 if mask is not None:202 qk = qk + mask203 qk = qk.float()204 205 w = F.softmax(qk, dim=-1).to(q.dtype)206 return (w @ v).permute(0, 2, 1, 3).flatten(start_dim=2), qk.detach()207 208class ResidualAttentionBlock(nn.Module):209 def __init__(self, n_state: int, n_head: int):210 super().__init__()211 212 self.attn = MultiHeadAttention(n_state, n_head)213 self.attn_ln = LayerNorm(n_state)214 215 n_mlp = n_state * 4216 self.mlp = nn.Sequential(217 Linear(n_state, n_mlp), nn.GELU(), Linear(n_mlp, n_state)218 )219 self.mlp_ln = LayerNorm(n_state)220 221 def forward(222 self,223 x: Tensor,224 mask: Optional[Tensor] = None,225 ):226 x = x + self.attn(self.attn_ln(x.contiguous()), mask=mask)[0]227 x = x + self.mlp(self.mlp_ln(x.contiguous()))228 return x229 230class AudioEncoder(nn.Module):231 def __init__(232 self, n_mels: int, n_ctx: int, n_state: int, n_head: int, n_layer: int233 ):234 super().__init__()235 self.conv1 = Conv1d(n_mels, n_state, kernel_size=3, padding=1)236 self.conv2 = Conv1d(n_state, n_state, kernel_size=3, stride=2, padding=1)237 self.positional_embedding = nn.Embedding(n_ctx, n_state)238 self.positional_embedding.requires_grad_(False)239 self.blocks: Iterable[ResidualAttentionBlock] = nn.ModuleList(240 [ResidualAttentionBlock(n_state, n_head) for _ in range(n_layer)]241 )242 self.avg_pooler = nn.AvgPool1d(2, stride=2)243 self.after_norm = LayerNorm(n_state)244 self.gradient_checkpointing = False245 246 def forward(self, x: Tensor, x_len: Tensor) -> Tuple[Tensor, Tensor]:247 T = x.size(-1)248 x = F.gelu(self.conv1(x))249 x = F.gelu(self.conv2(x))250 x = x.permute(0, 2, 1) # (B, T // 2, n_state)251 mask = make_non_pad_mask(x_len, T).unsqueeze(1) # (B, 1, T)252 mask = mask_to_bias(mask[:, :, (T + 1) % 2::2], x.dtype) # (B, 1, T // 2)253 x = (x + self.positional_embedding.weight[:x.shape[1], :]).to(x.dtype)254 for block in self.blocks:255 if self.gradient_checkpointing and self.training:256 x = torch.utils.checkpoint.checkpoint(block, x, mask.unsqueeze(1))257 else:258 x = block(x, mask.unsqueeze(1))259 x = x.permute(0, 2, 1)260 x = self.avg_pooler(x)261 x = x.permute(0, 2, 1)262 x_len = (x_len + 1) // 2 // 2263 x = self.after_norm(x.contiguous())264 return x, x_len265 266class Adaptor(nn.Module):267 def __init__(268 self,269 n_state: int = 1280,270 n_hidden: int = 3072,271 kernel_size: int = 7,272 stride: int = 4273 ):274 super().__init__()275 self.stride = stride276 if self.stride != -1:277 # print("self.stride: {}".format(self.stride))278 self.conv = Conv1d(n_state, n_state, kernel_size, stride, padding=1)279 self.linear1 = nn.Linear(n_state, 2048)280 self.relu = nn.ReLU()281 self.linear2 = nn.Linear(2048, n_hidden)282 self.gradient_checkpointing = False283 284 def forward(self, x: Tensor) -> Tuple[Tensor]:285 T = x.size(-1)286 if self.stride != -1:287 if self.gradient_checkpointing and self.training:288 x = torch.utils.checkpoint.checkpoint(self.conv, x.permute(0, 2, 1))289 x = x.permute(0, 2, 1)290 else:291 x = x.permute(0, 2, 1)292 x = F.gelu(self.conv(x))293 x = x.permute(0, 2, 1)294 if self.gradient_checkpointing and self.training:295 x = torch.utils.checkpoint.checkpoint(self.linear1, x)296 x = torch.utils.checkpoint.checkpoint(self.relu, x)297 x = torch.utils.checkpoint.checkpoint(self.linear2, x)298 else:299 x = self.linear1(x)300 x = self.relu(x)301 x = self.linear2(x)302 return x303 304class StepAudio2ForCausalLM(PreTrainedModel, GenerationMixin):305 config_class = StepAudio2Config306 main_input_name = "input_ids"307 # Important: Add this attribute to make HF recognize it as a model with generation capability308 # _keys_to_ignore_on_load_missing = ["lm_head.weight"]309 supports_gradient_checkpointing = True # 新增,声明支持gradient checkpointing310 311 def __init__(self, config: StepAudio2Config):312 super().__init__(config)313 if isinstance(config.torch_dtype, str):314 dtype = getattr(torch, config.torch_dtype)315 else:316 dtype = config.torch_dtype317 self.model = Qwen2Model(config.text_config)318 self.bf16 = dtype==torch.bfloat16319 self.encoder = AudioEncoder(320 config.audio_encoder_config.n_mels, config.audio_encoder_config.n_audio_ctx, config.audio_encoder_config.n_audio_state,321 config.audio_encoder_config.n_audio_head, config.audio_encoder_config.n_audio_layer322 )323 self.adapter = Adaptor(324 config.audio_encoder_config.n_audio_state, config.audio_encoder_config.llm_dim,325 config.audio_encoder_config.kernel_size, config.audio_encoder_config.adapter_stride326 )327 if self.bf16:328 self.encoder = self.encoder.bfloat16()329 self.adapter = self.adapter.bfloat16()330 self.lm_head = torch.nn.Linear(331 config.text_config.hidden_size,332 config.text_config.vocab_size,333 bias=False,334 dtype=dtype335 )336 self.post_init()337 338 def forward(339 self,340 input_ids=None,341 wavs=None,342 wav_lens=None,343 attention_mask=None,344 **kwargs345 ):346 hidden_states = self.model.embed_tokens(input_ids)347 if wavs is not None:348 if self.bf16:349 wavs = wavs.bfloat16()350 out, feat_lens = self.encoder(wavs, wav_lens)351 out = self.adapter(out)352 feat_lens = (feat_lens - 1) // 2 + 1353 insert_location = torch.nonzero(input_ids == 151688)354 insert_location[:,1] += 1355 for idx in range(len(insert_location)):356 i,s = insert_location[idx]357 hidden_states[i][s : s+feat_lens[idx]] = out[idx][:feat_lens[idx]]358 359 x = self.model(inputs_embeds=hidden_states, attention_mask=attention_mask)[0]360 logits = self.lm_head(x)361 return CausalLMOutputWithPast(362 logits=logits,363 past_key_values=None,364 hidden_states=None,365 attentions=None366 )367 368 def get_input_embeddings(self):369 """Return the model's input embeddings - required for GenerationMixin"""370 return self.model.embed_tokens371 372 def get_output_embeddings(self):373 """Return the model's output embeddings (LM head) - required for GenerationMixin"""374 return self.lm_head375 376 def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs):377 """Prepare inputs for generation - required for GenerationMixin"""378 # Keep the wavs and wav_lens from the initial call379 wavs = kwargs.get("wavs", None)380 wav_lens = kwargs.get("wav_lens", None)381 382 # For generation steps after the first, we don't need to process audio again383 # because the audio tokens have already been replaced in the input sequence384 if "past_key_values" in kwargs and kwargs["past_key_values"] is not None:385 # We're in a generation step, no need to process audio again386 return {387 "input_ids": input_ids,388 "attention_mask": attention_mask,389 "past_key_values": kwargs.get("past_key_values")390 }391 392 # First generation step, include audio processing393 return {394 "input_ids": input_ids,395 "attention_mask": attention_mask,396 "wavs": wavs,397 "wav_lens": wav_lens398 }399 400 def _reorder_cache(self, past_key_values, beam_idx):401 """Reorder the cache for beam search - required for GenerationMixin if using beam search"""402 # If you're not using past_key_values or beam search, this can be a simple pass-through403 # Otherwise implement according to your model's cache structure404 return past_key_values405 406 def _set_gradient_checkpointing(self, module, value=False):407 # For Qwen2Model408 if hasattr(self.model, 'gradient_checkpointing'):409 self.model.gradient_checkpointing = value410 411 # Add the missing _gradient_checkpointing_func method to Qwen2Model412 # This is what Qwen2Model tries to use when gradient_checkpointing=True413 if value and not hasattr(self.model, '_gradient_checkpointing_func'):414 def _gradient_checkpointing_func(module_to_run, *args, **kwargs):415 # This function wraps torch.utils.checkpoint.checkpoint416 # and is used by Qwen2Model to perform checkpointing417 return torch.utils.checkpoint.checkpoint(module_to_run, *args, **kwargs)418 419 self.model._gradient_checkpointing_func = _gradient_checkpointing_func420 421 # For custom encoder and adapter422 if hasattr(self.encoder, 'gradient_checkpointing'):423 self.encoder.gradient_checkpointing = value424 if hasattr(self.adapter, 'gradient_checkpointing'):425 self.adapter.gradient_checkpointing = value426 