#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
#           This file was automatically generated from src/transformers/models/vibevoice/modular_vibevoice.py.
#               Do NOT edit this file manually as any edits will be overwritten by the generation of
#             the file from the modular. If any change should be done, please apply the change to the
#                          modular_vibevoice.py file directly. One of our CI enforces this.
#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
# Copyright 2026 The Microsoft Team and The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import math
from dataclasses import dataclass

import torch
import torch.nn as nn
import torch.nn.functional as F

from ... import initialization as init
from ...activations import ACT2FN
from ...integrations import use_kernel_forward_from_hub
from ...modeling_outputs import BaseModelOutputWithPast, BaseModelOutputWithPooling
from ...modeling_utils import PreTrainedModel
from ...processing_utils import Unpack
from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, torch_compilable_check
from ..auto import AutoModel
from .configuration_vibevoice import VibeVoiceConfig
from .generation_vibevoice import VibeVoiceGenerationMixin


@auto_docstring(custom_intro="""Base class for VibeVoice outputs, with hidden states and attentions.""")
@dataclass
class VibeVoiceBaseModelOutputWithPast(BaseModelOutputWithPast):
    r"""
    audio_features (`torch.FloatTensor` of shape `(batch_size, sequence_length, feature_size)`, *optional*):
        Extracted audio features that can be used for conditioning the language model.
    """

    audio_features: torch.FloatTensor | None = None


@auto_docstring(custom_intro="""VibeVoice causal language model outputs.""")
@dataclass
class VibeVoiceCausalLMOutputWithPast(BaseModelOutputWithPast):
    r"""
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction).
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    audio_features (`torch.FloatTensor` of shape `(batch_size, sequence_length, feature_size)`, *optional*):
        Extracted audio features that can be used for conditioning the language model.
    """

    loss: torch.FloatTensor | None = None
    logits: torch.FloatTensor | None = None
    audio_features: torch.FloatTensor | None = None


@use_kernel_forward_from_hub("RMSNorm")
class VibeVoiceRMSNorm(nn.Module):
    def __init__(self, hidden_size, eps: float = 1e-6) -> None:
        """
        VibeVoiceRMSNorm is equivalent to T5LayerNorm
        """
        super().__init__()
        self.weight = nn.Parameter(torch.ones(hidden_size))
        self.variance_epsilon = eps

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        input_dtype = hidden_states.dtype
        hidden_states = hidden_states.to(torch.float32)
        variance = hidden_states.pow(2).mean(-1, keepdim=True)
        hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
        return self.weight * hidden_states.to(input_dtype)

    def extra_repr(self):
        return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"


class VibeVoiceDiffusionHeadSinusoidalEmbedding(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config
        self.freq = nn.Buffer(self.compute_default_freq(config), persistent=False)

    @staticmethod
    def compute_default_freq(config):
        dim = config.frequency_embedding_size // 2
        return torch.exp(-math.log(config.diffusion_max_period) * torch.arange(dim, dtype=torch.float32) / dim)

    def forward(self, timesteps):
        freq = timesteps[:, None].float() * self.freq[None].to(timesteps.device)
        embedding = torch.cat([torch.cos(freq), torch.sin(freq)], dim=-1)
        return embedding


class VibeVoiceDiffusionHeadMLP(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config
        self.activation_fn = ACT2FN[config.hidden_act]
        self.fc1 = nn.Linear(config.frequency_embedding_size, config.hidden_size, bias=False)
        self.fc2 = nn.Linear(config.hidden_size, config.hidden_size, bias=False)

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        hidden_states = self.fc1(hidden_states)
        hidden_states = self.activation_fn(hidden_states)
        hidden_states = self.fc2(hidden_states)
        return hidden_states


class VibeVoiceMLP(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config
        self.hidden_size = config.hidden_size
        self.intermediate_size = config.intermediate_size
        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias)
        self.act_fn = ACT2FN[config.hidden_act]

    def forward(self, x):
        down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
        return down_proj


class VibeVoiceDiffusionHeadAdaLayerNorm(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.num_chunks = 3
        self.ffn = VibeVoiceMLP(config)
        self.norm = VibeVoiceRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.act_fn = ACT2FN[config.hidden_act]
        self.linear = nn.Linear(config.hidden_size, config.hidden_size * self.num_chunks, bias=False)

    def forward(self, hidden_states, condition):
        shift_ffn, scale_ffn, gate_ffn = self.linear(self.act_fn(condition)).chunk(self.num_chunks, dim=-1)
        modulated_hidden_states = self.norm(hidden_states) * (1 + scale_ffn) + shift_ffn
        hidden_states = hidden_states + gate_ffn * self.ffn(modulated_hidden_states)
        return hidden_states


class VibeVoiceDiffusionHeadFinalLayer(nn.Module):
    def __init__(self, config, output_size):
        super().__init__()
        self.num_chunks = 2
        self.rms_norm_eps = config.rms_norm_eps
        self.linear_1 = nn.Linear(config.hidden_size, self.num_chunks * config.hidden_size, bias=False)
        self.act_fn = ACT2FN[config.hidden_act]
        self.linear_2 = nn.Linear(config.hidden_size, output_size, bias=False)

    def forward(self, hidden_states, condition):
        shift, scale = self.linear_1(self.act_fn(condition)).chunk(self.num_chunks, dim=-1)
        normed = F.rms_norm(hidden_states, (hidden_states.shape[-1],), eps=self.rms_norm_eps)
        hidden_states = normed * (1 + scale) + shift
        hidden_states = self.linear_2(hidden_states)
        return hidden_states


class VibeVoiceDiffusionHead(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.noisy_images_proj = nn.Linear(config.latent_size, config.hidden_size, bias=False)
        self.cond_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
        self.timestep_embedding = VibeVoiceDiffusionHeadSinusoidalEmbedding(config)
        self.timestep_proj = VibeVoiceDiffusionHeadMLP(config)
        self.layers = nn.ModuleList(
            [VibeVoiceDiffusionHeadAdaLayerNorm(config) for _ in range(config.num_hidden_layers)]
        )
        self.final_layer = VibeVoiceDiffusionHeadFinalLayer(config, output_size=config.latent_size)

    def forward(self, noisy_images, timesteps, condition):
        """
        Forward pass of the prediction head. Returns the predicted noise/velocity.
        """
        hidden_states = self.noisy_images_proj(noisy_images)
        embedded_timesteps = self.timestep_embedding(timesteps).to(timesteps.dtype)
        embedded_timesteps = self.timestep_proj(embedded_timesteps)
        condition = self.cond_proj(condition)
        condition = condition + embedded_timesteps
        for layer in self.layers:
            hidden_states = layer(hidden_states, condition)
        hidden_states = self.final_layer(hidden_states, condition)
        return hidden_states


class VibeVoiceMultiModalProjector(nn.Module):
    def __init__(self, input_dim, output_dim):
        super().__init__()
        self.linear_1 = nn.Linear(input_dim, output_dim)
        self.act = VibeVoiceRMSNorm(output_dim, eps=1e-6)
        self.linear_2 = nn.Linear(output_dim, output_dim)

    def forward(self, audio_features):
        hidden_states = self.linear_1(audio_features)
        hidden_states = self.act(hidden_states)
        hidden_states = self.linear_2(hidden_states)
        return hidden_states


@auto_docstring
class VibeVoicePreTrainedModel(PreTrainedModel):
    config: VibeVoiceConfig
    base_model_prefix = "model"
    input_modalities = ("audio", "text")
    supports_gradient_checkpointing = True
    _no_split_modules = ["VibeVoiceDiffusionHead"]
    _skip_keys_device_placement = ["past_key_values"]
    _supports_flash_attn = True
    _supports_sdpa = True
    _supports_flex_attn = True
    _supports_cache_class = True
    _supports_attention_backend = True
    _can_compile_fullgraph = True

    def _init_weights(self, module):
        super()._init_weights(module)
        if hasattr(module, "latent_scaling_factor"):
            init.constant_(module.latent_scaling_factor, 1.0)
        if hasattr(module, "latent_bias_factor"):
            init.constant_(module.latent_bias_factor, 0.0)
        if isinstance(module, VibeVoiceDiffusionHeadSinusoidalEmbedding):
            init.copy_(module.freq, module.compute_default_freq(module.config))


@auto_docstring(
    custom_intro="""
    The VibeVoice model which consists of audio tokenizers and an LLM backbone, without a language modeling head.
    """
)
class VibeVoiceModel(VibeVoicePreTrainedModel):
    def __init__(self, config):
        super().__init__(config)
        self.audio_tower = AutoModel.from_config(config.audio_config)
        self.language_model = AutoModel.from_config(config.text_config)
        self.multi_modal_projector = VibeVoiceMultiModalProjector(
            config.audio_config.hidden_size, config.text_config.hidden_size
        )
        self.semantic_tokenizer_encoder = AutoModel.from_config(config.semantic_model_config)
        self.semantic_connector = VibeVoiceMultiModalProjector(
            config.semantic_model_config.hidden_size, config.text_config.hidden_size
        )
        self.diffusion_head = VibeVoiceDiffusionHead(config.diffusion_head_config)
        self.latent_scaling_factor = nn.Parameter(torch.tensor(1.0))
        self.latent_bias_factor = nn.Parameter(torch.tensor(0.0))
        self.post_init()

    @can_return_tuple
    @auto_docstring(
        custom_intro="This method is used to get the audio embeddings (that replace placeholder audio tokens in the "
        "input sequence) and the acoustic features (used as diffusion target) from the input audio waveform."
    )
    def get_audio_features(
        self,
        input_values: torch.FloatTensor,
        padding_mask: torch.Tensor,
        **kwargs: Unpack[TransformersKwargs],
    ) -> tuple | BaseModelOutputWithPooling:
        r"""
        input_values (`torch.FloatTensor`):
            Float values of (normalized) audio waveform.
        padding_mask (`torch.Tensor` of shape `(batch_size, padded_audio_length)`):
            Padding mask to remove padded parts of audio.
        """

        # Acoustic tokenizer is not meant to be trainable (see p. 3 of https://huggingface.co/papers/2508.19205)
        with torch.no_grad():
            acoustic_features = self.audio_tower.encode(input_values, sample=True).latents
        acoustic_features += self.latent_bias_factor.to(acoustic_features.device)
        acoustic_features *= self.latent_scaling_factor.to(acoustic_features.device)

        # adjust padding mask according to tokenizer compression
        num_audio_tokens = torch.ceil(padding_mask.sum(dim=-1) / self.config.audio_config.hop_length).to(torch.int64)
        padding_mask = torch.arange(num_audio_tokens.max(), device=num_audio_tokens.device) < num_audio_tokens[:, None]

        pooler_output = self.multi_modal_projector(acoustic_features)
        return BaseModelOutputWithPooling(
            last_hidden_state=acoustic_features[padding_mask.to(acoustic_features.device)],
            pooler_output=pooler_output[padding_mask.to(pooler_output.device)],
        )

    def get_placeholder_mask(
        self, input_ids: torch.LongTensor, inputs_embeds: torch.FloatTensor, audio_features: torch.FloatTensor
    ):
        """
        Obtains multimodal placeholder mask from `input_ids` or `inputs_embeds`, and checks that the placeholder token count is
        equal to the length of multimodal features. If the lengths are different, an error is raised.
        """
        if input_ids is None:
            special_audio_mask = inputs_embeds == self.get_input_embeddings()(
                torch.full((), self.config.audio_token_id, dtype=torch.long, device=inputs_embeds.device)
            )
            special_audio_mask = special_audio_mask.all(-1)
        else:
            special_audio_mask = input_ids == self.config.audio_token_id

        n_audio_tokens = special_audio_mask.sum()
        n_audio_features = audio_features.shape[0]
        special_audio_mask = special_audio_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device)
        torch_compilable_check(
            inputs_embeds[special_audio_mask].numel() == audio_features.numel(),
            f"Audio features and audio tokens do not match, tokens: {n_audio_tokens}, features: {n_audio_features}",
        )
        return special_audio_mask

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor = None,
        inputs_embeds: torch.FloatTensor | None = None,
        input_values: torch.FloatTensor | None = None,
        padding_mask: torch.BoolTensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> tuple | BaseModelOutputWithPast:
        r"""
        padding_mask (`torch.Tensor` of shape `(batch_size, padded_audio_length)`):
            Padding mask to remove padded parts of audio.
        """
        if inputs_embeds is None:
            inputs_embeds = self.get_input_embeddings()(input_ids)

        audio_features = None
        if input_values is not None:
            audio_outputs = self.get_audio_features(input_values, padding_mask)
            audio_embeds = audio_outputs.pooler_output
            audio_features = audio_outputs.last_hidden_state

            # Replace text-audio token placeholders with audio embeddings
            special_audio_mask = self.get_placeholder_mask(
                input_ids, inputs_embeds=inputs_embeds, audio_features=audio_embeds
            )
            inputs_embeds = inputs_embeds.masked_scatter(special_audio_mask, audio_embeds.to(inputs_embeds.device))

        outputs = self.language_model(inputs_embeds=inputs_embeds, **kwargs)
        return VibeVoiceBaseModelOutputWithPast(audio_features=audio_features, **outputs)


@auto_docstring(
    custom_intro="""
    The VibeVoice model, which consists of a language model, audio tokenizers, connectors, and a diffusion head.
    """
)
class VibeVoiceForConditionalGeneration(VibeVoicePreTrainedModel, VibeVoiceGenerationMixin):
    _tied_weights_keys = {"lm_head.weight": "model.language_model.embed_tokens.weight"}
    _tp_plan = {"lm_head": "colwise_gather_output"}

    def __init__(self, config):
        super().__init__(config)
        self.model = VibeVoiceModel(config)
        self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
        self.post_init()

    @staticmethod
    def compute_diffusion_loss(
        diffusion_head: nn.Module,
        audio_features: torch.FloatTensor,
        condition_features: torch.FloatTensor,
        noise_scheduler: object,
        ddpm_batch_multiplier: int,
        num_diffusion_steps: int,
    ) -> torch.FloatTensor:
        audio_len, latent_size = audio_features.shape
        audio_features_repeated = audio_features.repeat_interleave(ddpm_batch_multiplier, dim=0)
        condition_features_repeated = condition_features.repeat_interleave(ddpm_batch_multiplier, dim=0)

        # Random noise and timesteps for diffusion training
        noise = torch.randn(
            (audio_len * ddpm_batch_multiplier, latent_size),
            device=audio_features.device,
            dtype=audio_features.dtype,
        )
        timesteps = torch.multinomial(
            torch.ones(num_diffusion_steps),
            audio_len * ddpm_batch_multiplier,
            replacement=True,
        ).to(audio_features.device)
        noisy_audio_features = noise_scheduler.add_noise(audio_features_repeated, noise, timesteps)

        # Predict noise/velocity using diffusion head
        model_output = diffusion_head(
            noisy_audio_features, timesteps.type_as(audio_features), condition_features_repeated
        )

        # Compute target (v_prediction parameterization)
        alpha_t = noise_scheduler.alpha_t.to(device=audio_features.device, dtype=audio_features.dtype)
        sigma_t = noise_scheduler.sigma_t.to(device=audio_features.device, dtype=audio_features.dtype)
        alpha_t = alpha_t[timesteps].flatten().unsqueeze(-1)
        sigma_t = sigma_t[timesteps].flatten().unsqueeze(-1)
        target = (alpha_t * noise - sigma_t * audio_features_repeated).to(model_output.device)

        diffusion_loss = torch.nn.functional.mse_loss(model_output.float(), target.float(), reduction="sum")
        if latent_size > 0 and ddpm_batch_multiplier > 0 and audio_len > 0:
            diffusion_loss = diffusion_loss / (latent_size * ddpm_batch_multiplier * audio_len)
        else:
            diffusion_loss = torch.full((), 0.0, dtype=diffusion_loss.dtype, device=diffusion_loss.device)
        return diffusion_loss

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor = None,
        inputs_embeds: torch.FloatTensor | None = None,
        labels: torch.LongTensor | None = None,
        logits_to_keep: int | slice = 0,
        input_values: torch.FloatTensor | None = None,
        padding_mask: torch.BoolTensor | None = None,
        acoustic_loss_mask: torch.BoolTensor | None = None,
        noise_scheduler: object | None = None,
        ddpm_batch_multiplier: int = 4,
        num_diffusion_steps: int = 10,
        **kwargs,
    ) -> tuple | VibeVoiceCausalLMOutputWithPast:
        r"""
        input_values (`torch.FloatTensor`, *optional*):
            Preprocessed audio waveform for voice cloning.
        padding_mask (`torch.BoolTensor`, *optional*):
            Masks indicating valid input frames.
        acoustic_loss_mask (`torch.BoolTensor`, *optional*):
            Mask to compute diffusion loss only on specific acoustic tokens.
        noise_scheduler (*optional*):
            Needed for training to compute noise targets for the diffusion loss. By default, uses the noise scheduler
            configuration specified in the model's `generation_config`.
        ddpm_batch_multiplier (`int`, *optional*, defaults to 4):
            For training, number of noise samples to generate per audio token for diffusion loss computation, which can
            help stabilize training.
        num_diffusion_steps (`int`, *optional*, defaults to 10):
            For training, the number of diffusion steps to use. Defaults to 10 if not provided.
        """

        outputs = self.model(
            input_ids=input_ids,
            inputs_embeds=inputs_embeds,
            input_values=input_values,
            padding_mask=padding_mask,
            **kwargs,
        )

        hidden_states = outputs.last_hidden_state
        slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
        logits = self.lm_head(hidden_states[:, slice_indices, :])

        loss = None
        if labels is not None:
            loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs)

        # `diffusion_loss_weight` weights the two training losses against each other, and does not affect the architecture
        # trf-ignore: TRF041 config.diffusion_loss_weight
        if (
            acoustic_loss_mask is not None
            and outputs.audio_features is not None
            and self.config.diffusion_loss_weight > 0.0
        ):
            if noise_scheduler is None:
                # use helper from VibeVoiceGenerationMixin
                noise_scheduler = self._build_default_noise_scheduler(self.generation_config)

            audio_features = outputs.audio_features
            condition_features = hidden_states[acoustic_loss_mask.to(hidden_states.device)].to(audio_features.device)
            diffusion_loss = self.compute_diffusion_loss(
                self.model.diffusion_head,
                audio_features,
                condition_features,
                noise_scheduler,
                ddpm_batch_multiplier,
                num_diffusion_steps,
            )
            loss = (
                1 - self.config.diffusion_loss_weight
            ) * loss + self.config.diffusion_loss_weight * diffusion_loss.to(loss.device)

        return VibeVoiceCausalLMOutputWithPast(loss=loss, logits=logits, **outputs)


__all__ = ["VibeVoiceForConditionalGeneration", "VibeVoicePreTrainedModel", "VibeVoiceModel"]
