CoolFace
Apppublic

vincenthugging/MOSS-TTSD-Enhanced

sourceHugging Faceupdated 1y agoView on Hugging Face
6likes
modeling_asteroid.py446 linesDownload Raw Back to root
1import torch2import torch.nn as nn3from dataclasses import dataclass4from transformers.utils import ModelOutput5from transformers.cache_utils import Cache6from typing import Optional, List, Tuple, Union7from transformers.loss.loss_utils import ForCausalLMLoss8from transformers.generation.streamers import BaseStreamer9from transformers.modeling_outputs import BaseModelOutputWithPast10from transformers.generation.configuration_utils import GenerationConfig11from transformers.generation.stopping_criteria import StoppingCriteriaList12from transformers import PreTrainedModel, GenerationMixin, Qwen3Config, Qwen3Model13from transformers.generation.logits_process import LogitsProcessorList, RepetitionPenaltyLogitsProcessor, TopKLogitsWarper, TopPLogitsWarper, TemperatureLogitsWarper14try:15    from liger_kernel.transformers.model.loss_utils import LigerForCausalLMLoss16    LIGER_AVAILABLE = True17except ImportError:18    print("Warning: liger_kernel not available, using standard CrossEntropyLoss")19    LigerForCausalLMLoss = None20    LIGER_AVAILABLE = False21 22 23class AsteroidTTSConfig(Qwen3Config):24    def __init__(self, 25                channels = 8,26                speech_pad_token = 1024,27                speech_vocab_size = 1025,28                speech_token_range = [],29                **kwargs):30        super().__init__(**kwargs)31        self.channels = channels32        self.speech_pad_token = speech_pad_token33        self.speech_vocab_size = speech_vocab_size34        self.speech_token_range = speech_token_range35        36 37@dataclass38class AsteroidTTSOutputWithPast(ModelOutput):39    loss: Optional[torch.FloatTensor] = None40    logits: torch.FloatTensor = None41    loss_all: Optional[Tuple[torch.FloatTensor]] = None42    logits_all: Optional[Tuple[torch.FloatTensor]] = None43    past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None44    hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None45    attentions: Optional[Tuple[torch.FloatTensor, ...]] = None46    47 48@dataclass49class GenerateDecoderOnlyOutput(ModelOutput):50    sequences: torch.LongTensor = None51    scores: Optional[Tuple[torch.FloatTensor]] = None52    logits: Optional[Tuple[torch.FloatTensor]] = None53    attentions: Optional[Tuple[Tuple[torch.FloatTensor]]] = None54    hidden_states: Optional[Tuple[Tuple[torch.FloatTensor]]] = None55    past_key_values: Optional[Tuple[Tuple[Tuple[torch.FloatTensor]]]] = None56 57 58class CustomMixin(GenerationMixin):59    def _sample(60        self,61        input_ids: torch.LongTensor,62        logits_processor: LogitsProcessorList,63        stopping_criteria: StoppingCriteriaList,64        generation_config: GenerationConfig,65        synced_gpus: bool,66        streamer: Optional["BaseStreamer"],67        **model_kwargs,68    ) -> Union[GenerateDecoderOnlyOutput, torch.LongTensor]:69        # Extract configuration parameters70        speech_pad_idx = self.config.speech_pad_token71        72        eos_token_id = generation_config.eos_token_id73        output_attentions = generation_config.output_attentions74        output_hidden_states = generation_config.output_hidden_states75        output_scores = generation_config.output_scores76        output_logits = generation_config.output_logits77        return_dict_in_generate = generation_config.return_dict_in_generate78        max_length = generation_config.max_length79        has_eos_stopping_criteria = any(hasattr(criteria, "eos_token_id") for criteria in stopping_criteria)80        do_sample = generation_config.do_sample81 82        # Initialize output tuples83        scores = () if (return_dict_in_generate and output_scores) else None84        raw_logits = () if (return_dict_in_generate and output_logits) else None85        decoder_attentions = () if (return_dict_in_generate and output_attentions) else None86        decoder_hidden_states = () if (return_dict_in_generate and output_hidden_states) else None87 88        # Initialize tracking variables89        batch_size, cur_len, channels = input_ids.shape  # channels = 890        this_peer_finished = False91        unfinished_sequences = torch.ones(batch_size, dtype=torch.long, device=input_ids.device)92        needs_additional_steps = -1 * torch.ones(batch_size, dtype=torch.long, device=input_ids.device)93        tf_inputs = input_ids[:]94        input_ids = input_ids[:, :-(channels - 1)]95        cur_len = input_ids.shape[1]96        model_kwargs["attention_mask"] = model_kwargs["attention_mask"][:, :-(channels - 1)]97        base_length = input_ids.shape[1]98        model_kwargs = self._get_initial_cache_position(cur_len, input_ids.device, model_kwargs)99 100        # Define logits processor101        if generation_config.do_samples is not None:102            do_samples = generation_config.do_samples103            realprocessor = [LogitsProcessorList() for _ in range(channels)]104            for i, layer_config in enumerate(generation_config.layers):105                if layer_config.get("repetition_penalty") is not None:106                    realprocessor[i].append(RepetitionPenaltyLogitsProcessor(penalty=layer_config.get("repetition_penalty")))107                if layer_config.get("temperature") is not None: 108                    realprocessor[i].append(TemperatureLogitsWarper(temperature=layer_config.get("temperature")))109                if layer_config.get("top_k") is not None:110                    realprocessor[i].append(TopKLogitsWarper(top_k=layer_config.get("top_k")))111                if layer_config.get("top_p") is not None:112                    realprocessor[i].append(TopPLogitsWarper(top_p=layer_config.get("top_p")))113        else:114            do_samples = [do_sample for _ in range(channels)]115            realprocessor = [logits_processor for _ in range(channels)]116        while self._has_unfinished_sequences(this_peer_finished, synced_gpus, device=input_ids.device):117            # Prepare model inputs118            model_inputs = self.prepare_inputs_for_generation(input_ids, **model_kwargs)119            model_inputs.update({"output_attentions": output_attentions} if output_attentions else {})120            model_inputs.update({"output_hidden_states": output_hidden_states} if output_hidden_states else {})121            # Forward pass122            outputs = self(**model_inputs, return_dict=True)123            model_kwargs = self._update_model_kwargs_for_generation(outputs, model_kwargs)124 125            if synced_gpus and this_peer_finished:126                continue127 128            # Get next token logits129            next_token_logits = [logits[:, -1, :].clone().float().to(input_ids.device) for logits in outputs.logits_all]130            for i, channel_logits in enumerate(next_token_logits):131                if i != 0 and input_ids.shape[1] + 1 > tf_inputs.shape[1] - 7 + i: 132                    channel_logits[:, 1024] = - torch.inf133                if i == 0 and input_ids.shape[1] + 1 <= tf_inputs.shape[1]: 134                    channel_logits[:, 152694] = - torch.inf135            next_token_scores = [realprocessor[i](input_ids[..., i], logits) for i, logits in enumerate(next_token_logits)]136            # Generate next tokens137            next_tokens = []138            for i, channel_score in enumerate(next_token_scores):139                if do_samples[i]:140                    # 添加数值稳定性保护141                    # 检查并处理异常值142                    if torch.isnan(channel_score).any() or torch.isinf(channel_score).any():143                        print(f"⚠️ 检测到异常值,使用argmax采样")144                        channel_ntk = torch.argmax(channel_score, dim=-1)145                    else:146                        # 数值稳定的softmax计算147                        channel_score_stable = channel_score - torch.max(channel_score, dim=-1, keepdim=True)[0]148                        probs = nn.functional.softmax(channel_score_stable, dim=-1)149                        150                        # 确保概率值有效151                        probs = torch.clamp(probs, min=1e-8, max=1.0)152                        probs = probs / probs.sum(dim=-1, keepdim=True)  # 重新归一化153                        154                        channel_ntk = torch.multinomial(probs, num_samples=1).squeeze(1)155                elif not do_samples[i]:156                    channel_ntk = torch.argmax(channel_score, dim=-1)157                next_tokens.append(channel_ntk)158            next_tokens = torch.stack(next_tokens, dim=-1)  # [batch_size, channels]159            # Additional steps logic160            indices = (~self.is_speech_token(next_tokens[:, 0])) & (needs_additional_steps < 0)161            needs_additional_steps[indices] = channels - 1  # For 8 channels, need 7 steps162            163            if input_ids.shape[1] + 1 <= tf_inputs.shape[1]:164                i = input_ids.shape[1] + 1 - base_length165                next_tokens[:, i:] = tf_inputs[:, input_ids.shape[1], i:]166            167            # Replace tokens in additional steps168            mask = (needs_additional_steps > 0) & (needs_additional_steps < 7)169            if mask.any().item():170                next_tokens[mask, 0] = self.config.eos_token_id171                for i in range(1, channels):172                    mask_i = mask & (needs_additional_steps < channels - i)173                    next_tokens[mask_i, i] = speech_pad_idx174            175            if has_eos_stopping_criteria:176                for i in range(channels):177                    pddp = self.config.eos_token_id if i == 0 else speech_pad_idx178                    next_tokens[:, i] = next_tokens[:, i] * unfinished_sequences + pddp * (1 - unfinished_sequences)179                    180            input_ids = torch.cat([input_ids, next_tokens[:, None, :]], dim=1)181            if streamer is not None:182                streamer.put(next_tokens[:, 0].cpu())183            184            # Update unfinished_sequences185            needs_additional_steps = torch.where(needs_additional_steps > 0, needs_additional_steps - 1, needs_additional_steps)186            stopping = stopping_criteria(input_ids[..., 0], scores) | (needs_additional_steps == 0)187            unfinished_sequences = unfinished_sequences & ~stopping188            unfinished_sequences = unfinished_sequences | (needs_additional_steps > 0)189            this_peer_finished = unfinished_sequences.max() == 0190 191            if return_dict_in_generate:192                if output_scores:193                    scores += (next_token_scores,)194                if output_logits:195                    raw_logits += (next_token_logits,)196                if output_attentions:197                    decoder_attentions += (outputs.attentions,)198                if output_hidden_states:199                    decoder_hidden_states += (outputs.hidden_states,)200 201            cur_len += 1202            del outputs203            204        if streamer is not None:205            streamer.end()206 207        if return_dict_in_generate:208            return GenerateDecoderOnlyOutput(209                sequences=input_ids,210                scores=scores,211                logits=raw_logits,212                attentions=decoder_attentions,213                hidden_states=decoder_hidden_states,214                past_key_values=model_kwargs.get("past_key_values"),215            )216        else:217            return input_ids218    219    220class AsteroidTTSPretrainedModel(PreTrainedModel):221    config_class = AsteroidTTSConfig222    base_model_prefix = "model"223    supports_gradient_checkpointing = True224    _no_split_modules = ["Qwen3DecoderLayer"]225    _skip_keys_device_placement = ["past_key_values"]226    _supports_flash_attn_2 = True227    _supports_sdpa = True228    _supports_flex_attn = True229    _supports_cache_class = True230    _supports_quantized_cache = True231    _supports_static_cache = True232    _supports_attention_backend = True233 234 235class AsteroidTTSModel(AsteroidTTSPretrainedModel):236    def __init__(self, config: AsteroidTTSConfig):237        super().__init__(config)238        self.text_pad_idx = config.pad_token_id239        self.speech_pad_idx = config.speech_pad_token240        self.embedding_list = nn.ModuleList([])241        self.embedding_list.append(nn.Embedding(config.vocab_size, config.hidden_size, self.text_pad_idx))242        # Channels 1 to channels-1: Speech tokens only243        for _ in range(1, config.channels):244            self.embedding_list.append(nn.Embedding(config.speech_vocab_size, config.hidden_size, self.speech_pad_idx))245 246        self.language_model = Qwen3Model(config)247        self.post_init()248 249    def get_input_embeddings(self):250        return self.embedding_list[0]251 252    def set_input_embeddings(self, value: nn.Embedding):253        self.embedding_list[0] = value254 255    def _prepare_multi_modal_inputs(self, input_ids: torch.LongTensor) -> torch.FloatTensor:256        """257        Prepares multi-modal embeddings from input_ids of shape (batch_size, channels, sequence_length).258        For channel 0: text + speech tokens, for channels 1 to channels-1: speech tokens padded with speech_pad_token.259        """260        batch_size, seq_length, channels = input_ids.shape261        if channels != self.config.channels:262            raise ValueError(f"Expected {self.config.channels} channels, got {channels}")263        264        inputs_embeds = torch.zeros(batch_size, seq_length, self.config.hidden_size, device=input_ids.device, dtype=self.embedding_list[0].weight.dtype)265        for i in range(channels):266            embed_layer = self.embedding_list[i]267            channel_input = input_ids[...,i]268            inputs_embeds += embed_layer(channel_input)269 270        return inputs_embeds271 272    def forward(273        self,274        input_ids: torch.LongTensor = None,  # Shape: (batch_size, channels, sequence_length)275        attention_mask: Optional[torch.Tensor] = None,276        position_ids: Optional[torch.LongTensor] = None,277        past_key_values: Optional[List[torch.FloatTensor]] = None,278        inputs_embeds: Optional[torch.FloatTensor] = None,279        use_cache: Optional[bool] = None,280        output_attentions: Optional[bool] = None,281        output_hidden_states: Optional[bool] = None,282        return_dict: Optional[bool] = None,283        cache_position: Optional[torch.LongTensor] = None,284        **kwargs,285    ) -> Union[Tuple, BaseModelOutputWithPast]:286 287        if (input_ids is None) ^ (inputs_embeds is not None):288            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")289 290        if input_ids is not None:291            inputs_embeds = self._prepare_multi_modal_inputs(input_ids)292 293        outputs = self.language_model(294            input_ids=None,295            attention_mask=attention_mask,296            position_ids=position_ids,297            past_key_values=past_key_values,298            inputs_embeds=inputs_embeds,299            use_cache=use_cache,300            output_attentions=output_attentions,301            output_hidden_states=output_hidden_states,302            return_dict=return_dict,303            cache_position=cache_position,304        )305        return outputs306    307    308class AsteroidTTSInstruct(AsteroidTTSPretrainedModel, CustomMixin):309    _tied_weights_keys = []310    _tp_plan = {"lm_head": "colwise_rep"}311    _pp_plan = {"lm_head": (["hidden_states"], ["logits"])}312 313    def __init__(self, config: AsteroidTTSConfig):314        super().__init__(config)315        self.model = AsteroidTTSModel(config)316        self.channels = config.channels317        self.weights = [1 for _ in range(self.channels)]318        self._tied_weights_keys = [f"lm_heads.{i}.weight" for i in range(self.channels)]319        self.vocab_size = config.vocab_size320        self.lm_heads = nn.ModuleList([])321        self.lm_heads.append(nn.Linear(config.hidden_size, config.vocab_size, bias=False))322        for _ in range(1, config.channels):323            self.lm_heads.append(nn.Linear(config.hidden_size, config.speech_vocab_size, bias=False))324        self.post_init()325 326    def get_input_embeddings(self):327        return self.model.embedding_list[0]328    329    def can_generate(self):330        return True331    332    def is_speech_token(self, tokens):333        return (tokens >= self.config.speech_token_range[0]) & (tokens < self.config.speech_token_range[1])334    335    def tie_weights(self):336        for i in range(self.config.channels):337            self._tie_or_clone_weights(self.lm_heads[i], self.model.embedding_list[i])338 339    def set_input_embeddings(self, value):340        self.model.embedding_list[0] = value341 342    def get_output_embeddings(self):343        return self.lm_heads[0]344 345    def set_output_embeddings(self, new_embeddings):346        self.lm_heads[0] = new_embeddings347 348    def set_decoder(self, decoder):349        self.model = decoder350 351    def get_decoder(self):352        return self.model353    354    def set_weights(self, weights):355        self.weights = weights356 357    def forward(358        self,359        input_ids: torch.LongTensor = None,360        attention_mask: Optional[torch.Tensor] = None,361        position_ids: Optional[torch.LongTensor] = None,362        past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None,363        inputs_embeds: Optional[torch.FloatTensor] = None,364        labels: Optional[torch.LongTensor] = None,365        use_cache: Optional[bool] = None,366        output_attentions: Optional[bool] = None,367        output_hidden_states: Optional[bool] = None,368        return_dict: Optional[bool] = None,369        cache_position: Optional[torch.LongTensor] = None,370        skip_logits: Optional[bool] = None,371        **kwargs,372    ) -> Union[Tuple, AsteroidTTSOutputWithPast]:373        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions374        output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states375        return_dict = return_dict if return_dict is not None else self.config.use_return_dict376 377        skip_logits = skip_logits if skip_logits is not None else (self.training and labels is not None)378        if skip_logits and labels is None:379            skip_logits = False380 381        # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)382        outputs = self.model(383            input_ids=input_ids,384            attention_mask=attention_mask,385            position_ids=position_ids,386            past_key_values=past_key_values,387            inputs_embeds=inputs_embeds,388            use_cache=use_cache,389            output_attentions=output_attentions,390            output_hidden_states=output_hidden_states,391            return_dict=return_dict,392            cache_position=cache_position,393            **kwargs,394        )395 396        hidden_states = outputs[0]397 398        logits_all = None399        loss_all = None400        total_loss = None401        402        if labels is not None:403            device = input_ids.device if input_ids is not None else inputs_embeds.device404            loss_all = torch.empty(self.channels, device=device)405            logits_list = []406            407            for i in range(self.config.channels):408                vocab_size = self.config.vocab_size if i == 0 else self.config.speech_vocab_size409                if skip_logits and LIGER_AVAILABLE:410                    loss_all[i] = LigerForCausalLMLoss(411                        hidden_states=hidden_states,412                        lm_head_weight=self.lm_heads[i].weight,413                        labels=labels[..., i],414                        hidden_size=self.config.hidden_size,415                        **kwargs416                    )417                else:418                    logits = self.lm_heads[i](hidden_states)419                    loss_all[i] = ForCausalLMLoss(logits, labels[..., i], vocab_size)420                    logits_list.append(logits)421 422            if not skip_logits:423                logits_all = tuple(logits_list)424 425            total_weight = sum(self.weights)426            normalized_weights = [w / total_weight for w in self.weights]427            428            total_loss = 0429            for w, loss in zip(normalized_weights, loss_all):430                total_loss += w * loss431        else:432            logits_all = [lm_head(hidden_states) for lm_head in self.lm_heads]433 434        if not return_dict:435            output = (logits_all,) + outputs[1:]436            return (total_loss, loss_all, ) + output if total_loss is not None else output437 438        return AsteroidTTSOutputWithPast(439            loss=total_loss,440            logits=logits_all[0] if logits_all is not None else None,441            loss_all=loss_all,442            logits_all=logits_all,443            past_key_values=outputs.past_key_values,444            hidden_states=outputs.hidden_states,445            attentions=outputs.attentions,446        )