# 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.

from huggingface_hub.dataclasses import strict

from ...configuration_utils import PreTrainedConfig
from ...utils import auto_docstring
from ..auto import CONFIG_MAPPING, AutoConfig


@auto_docstring(checkpoint="vibevoice/VibeVoice-1.5B-hf")
@strict
class VibeVoiceDiffusionHeadConfig(PreTrainedConfig):
    r"""
    latent_size (`int`, *optional*, defaults to 64):
        Dimensionality of the acoustic latents the head denoises.
    frequency_embedding_size (`int`, *optional*, defaults to 256):
        The size of the sinusoidal frequency embedding for timestep encoding in the diffusion head.
    diffusion_max_period (`int`, *optional*, defaults to 10000):
        The maximum period for the sinusoidal frequency embedding in the diffusion head.
    """

    hidden_size: int = 1536
    latent_size: int = 64
    num_hidden_layers: int = 4
    intermediate_size: int = 4608
    rms_norm_eps: float = 1e-5
    hidden_act: str = "silu"
    frequency_embedding_size: int = 256
    diffusion_max_period: int = 10000
    mlp_bias: bool = False


@auto_docstring(checkpoint="vibevoice/VibeVoice-1.5B-hf")
@strict
class VibeVoiceConfig(PreTrainedConfig):
    r"""
    semantic_model_config (`Union[AutoConfig, dict]`, *optional*):
        The config object or dictionary of the semantic tokenizer encoder. This tokenizer extracts semantic features from audio.
    diffusion_head_config (`Union[VibeVoiceDiffusionHeadConfig, dict]`, *optional*):
        The config object or dictionary of the diffusion head used to synthesize acoustic latents.
    audio_bos_token_id (`int`, *optional*, defaults to 151652):
        The token ID indicating the start of audio tokens.
    audio_eos_token_id (`int`, *optional*, defaults to 151653):
        The token ID indicating the end of audio tokens.
    diffusion_loss_weight (`float`, *optional*, defaults to 0.5):
        The weight of the diffusion loss in the overall loss computation. The cross entropy loss for the language
        modeling head is weighted by `(1 - diffusion_loss_weight)`.

    ```python
    >>> from transformers import VibeVoiceForConditionalGeneration, VibeVoiceConfig

    >>> # Initializing a VibeVoice configuration
    >>> configuration = VibeVoiceConfig()

    >>> # Initializing a 1.5B model with random weights
    >>> model = VibeVoiceForConditionalGeneration(configuration)

    >>> # Accessing the model configuration
    >>> configuration = model.config
    ```"""

    model_type = "vibevoice"
    sub_configs = {
        "audio_config": AutoConfig,
        "semantic_model_config": AutoConfig,
        "text_config": AutoConfig,
        "diffusion_head_config": VibeVoiceDiffusionHeadConfig,
    }

    audio_config: dict | PreTrainedConfig | None = None
    semantic_model_config: dict | PreTrainedConfig | None = None
    text_config: dict | PreTrainedConfig | None = None
    diffusion_head_config: dict | PreTrainedConfig | None = None
    pad_token_id: int = 151643
    eos_token_id: int = 151643
    audio_bos_token_id: int = 151652
    audio_eos_token_id: int = 151653
    audio_token_id: int = 151654
    diffusion_loss_weight: float = 0.5

    def __post_init__(self, **kwargs):
        if isinstance(self.audio_config, dict):
            self.audio_config["model_type"] = self.audio_config.get("model_type", "vibevoice_acoustic_tokenizer")
            self.audio_config = CONFIG_MAPPING[self.audio_config["model_type"]](**self.audio_config)
        elif self.audio_config is None:
            self.audio_config = CONFIG_MAPPING["vibevoice_acoustic_tokenizer"]()

        if isinstance(self.semantic_model_config, dict):
            self.semantic_model_config["model_type"] = self.semantic_model_config.get(
                "model_type", "vibevoice_acoustic_tokenizer_encoder"
            )
            self.semantic_model_config = CONFIG_MAPPING[self.semantic_model_config["model_type"]](
                **self.semantic_model_config
            )
        elif self.semantic_model_config is None:
            self.semantic_model_config = CONFIG_MAPPING["vibevoice_acoustic_tokenizer_encoder"](hidden_size=128)

        if isinstance(self.text_config, dict):
            self.text_config["model_type"] = self.text_config.get("model_type", "qwen2")
            self.text_config = CONFIG_MAPPING[self.text_config["model_type"]](**self.text_config)
        elif self.text_config is None:
            self.text_config = CONFIG_MAPPING["qwen2"]()

        if isinstance(self.diffusion_head_config, dict):
            self.diffusion_head_config = VibeVoiceDiffusionHeadConfig(**self.diffusion_head_config)
        elif self.diffusion_head_config is None:
            self.diffusion_head_config = VibeVoiceDiffusionHeadConfig(
                hidden_size=self.text_config.hidden_size, latent_size=self.audio_config.hidden_size
            )

        self.vocab_size = self.text_config.vocab_size
        self.tie_word_embeddings = getattr(self.text_config, "tie_word_embeddings", False)
        super().__post_init__(**kwargs)

    def validate_architecture(self):
        """Part of `@strict`-powered validation. Validates the architecture of the config."""
        if self.diffusion_head_config.hidden_size != self.text_config.hidden_size:
            raise ValueError(
                f"`diffusion_head_config.hidden_size` ({self.diffusion_head_config.hidden_size}) must match "
                f"`text_config.hidden_size` ({self.text_config.hidden_size})."
            )
        if self.diffusion_head_config.latent_size != self.audio_config.hidden_size:
            raise ValueError(
                f"`diffusion_head_config.latent_size` ({self.diffusion_head_config.latent_size}) must match "
                f"`audio_config.hidden_size` ({self.audio_config.hidden_size})."
            )


__all__ = ["VibeVoiceConfig", "VibeVoiceDiffusionHeadConfig"]
