fix(llama.py): 适配cuda13.x及更新的库版本
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user