diff --git a/soulxsinger/models/modules/llama.py b/soulxsinger/models/modules/llama.py index f9fee9c..aaee789 100644 --- a/soulxsinger/models/modules/llama.py +++ b/soulxsinger/models/modules/llama.py @@ -56,7 +56,7 @@ class LlamaNARDecoderLayer(LlamaDecoderLayer): config.hidden_size, eps=config.rms_norm_eps, dim_cond=config.hidden_size ) - # add `cond` in forward function + # 修改点:添加 position_embeddings 参数和 **kwargs def forward( self, hidden_states: torch.Tensor, @@ -66,38 +66,39 @@ class LlamaNARDecoderLayer(LlamaDecoderLayer): past_key_value: Optional[Tuple[torch.Tensor]] = None, output_attentions: Optional[bool] = False, use_cache: Optional[bool] = False, + position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + **kwargs, ) -> Tuple[ torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]] ]: - """ - Args: - hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)` - attention_mask (`torch.FloatTensor`, *optional*): attention mask of size - `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values. - output_attentions (`bool`, *optional*): - Whether or not to return the attentions tensors of all attention layers. See `attentions` under - returned tensors for more detail. - use_cache (`bool`, *optional*): - If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding - (see `past_key_values`). - past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states - """ - residual = hidden_states hidden_states = self.input_layernorm( hidden_states, cond_embedding=cond_embedding ) - # Self Attention - hidden_states, self_attn_weights, present_key_value = self.self_attn( + # [MODIFIED] 兼容新版 transformers 的返回值解包逻辑 + # 因为新版在 output_attentions=False 时可能只返回 (hidden_states, present_key_value) + attn_outputs = self.self_attn( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=output_attentions, use_cache=use_cache, + position_embeddings=position_embeddings, ) + + hidden_states = attn_outputs[0] + + # 动态处理权重和 KV 缓存的解包 + if output_attentions: + self_attn_weights = attn_outputs[1] + present_key_value = attn_outputs[2] if use_cache else None + else: + self_attn_weights = None + present_key_value = attn_outputs[1] if use_cache else None + hidden_states = residual + hidden_states # Fully Connected @@ -129,7 +130,15 @@ class DiffLlama(LlamaModel): dropout=0.1, ffn_dropout=0.1, attention_dropout=0.0, - config=LlamaConfig(0, 256, 1024, 1, 1), + config=LlamaConfig( + vocab_size=0, + hidden_size=256, + intermediate_size=1024, + num_hidden_layers=1, + num_attention_heads=4, # 确保维度对齐 (256/4=64) + num_key_value_heads=4, + attn_implementation="eager", + ), ): super().__init__(config) @@ -293,6 +302,13 @@ class DiffLlama(LlamaModel): else: position_ids = position_ids.view(-1, seq_length).long() + # ============================================================== + # [NEW] 修复位置:为适配新版 transformers 生成 position_embeddings + # ============================================================== + # self.rotary_emb 继承自 LlamaModel + position_embeddings = self.rotary_emb(hidden_states if 'hidden_states' in locals() else inputs_embeds, position_ids) + # ============================================================== + # embed positions if attention_mask is None: attention_mask = torch.ones( @@ -300,6 +316,9 @@ class DiffLlama(LlamaModel): dtype=torch.bool, device=inputs_embeds.device, ) + + # 处理 attention_mask 格式 + # 注意:某些新版本 transformers 此处返回值结构有变,如报错请检查此函数 attention_mask = self._prepare_decoder_attention_mask( attention_mask, (batch_size, seq_length), @@ -329,23 +348,11 @@ class DiffLlama(LlamaModel): ) if self.gradient_checkpointing and self.training: + # 原始代码此处有 NotImplementedError,如果需要开启梯度检查点, + # 也需要将 position_embeddings 传入 custom_forward raise NotImplementedError - - def create_custom_forward(module): - def custom_forward(*inputs): - # None for past_key_value - return module(*inputs, output_attentions, None) - - return custom_forward - - layer_outputs = torch.utils.checkpoint.checkpoint( - create_custom_forward(decoder_layer), - hidden_states, - attention_mask, - position_ids, - None, - ) else: + # [MODIFIED] 在调用 decoder_layer 时传入 position_embeddings layer_outputs = decoder_layer( hidden_states, attention_mask=attention_mask, @@ -354,6 +361,7 @@ class DiffLlama(LlamaModel): output_attentions=output_attentions, use_cache=use_cache, cond_embedding=diffusion_step, + position_embeddings=position_embeddings, # 新增参数 ) hidden_states = layer_outputs[0] @@ -375,18 +383,10 @@ class DiffLlama(LlamaModel): hidden_states = self.mel_out_mlp(hidden_states) - # if not return_dict: - # return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None) - # return BaseModelOutputWithPast( - # last_hidden_state=hidden_states, - # past_key_values=next_cache, - # hidden_states=all_hidden_states, - # attentions=all_self_attns, - # ) if return_dict: return { "output": hidden_states, "hidden_states": all_layer_hidden_states, } - return hidden_states + return hidden_states \ No newline at end of file