CoolFace
Modelpublic

Firworks/Step-Audio-R1-nvfp4

sourceHugging Faceapache-2.0updated 10mo agoView on Hugging Face
2likes17downloads
modeling_step_audio_2.py426 linesDownload Raw Back to root
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