fix(llama.py): 适配cuda13.x及更新的库版本

This commit is contained in:
mei
2026-04-19 08:21:01 +08:00
parent 2c3ace3db4
commit 7a71e78b11
+41 -41
View File
@@ -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,14 +383,6 @@ 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,