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
|
config.hidden_size, eps=config.rms_norm_eps, dim_cond=config.hidden_size
|
||||||
)
|
)
|
||||||
|
|
||||||
# add `cond` in forward function
|
# 修改点:添加 position_embeddings 参数和 **kwargs
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -66,38 +66,39 @@ class LlamaNARDecoderLayer(LlamaDecoderLayer):
|
|||||||
past_key_value: Optional[Tuple[torch.Tensor]] = None,
|
past_key_value: Optional[Tuple[torch.Tensor]] = None,
|
||||||
output_attentions: Optional[bool] = False,
|
output_attentions: Optional[bool] = False,
|
||||||
use_cache: Optional[bool] = False,
|
use_cache: Optional[bool] = False,
|
||||||
|
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
|
**kwargs,
|
||||||
) -> Tuple[
|
) -> Tuple[
|
||||||
torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]
|
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
|
residual = hidden_states
|
||||||
|
|
||||||
hidden_states = self.input_layernorm(
|
hidden_states = self.input_layernorm(
|
||||||
hidden_states, cond_embedding=cond_embedding
|
hidden_states, cond_embedding=cond_embedding
|
||||||
)
|
)
|
||||||
|
|
||||||
# Self Attention
|
# [MODIFIED] 兼容新版 transformers 的返回值解包逻辑
|
||||||
hidden_states, self_attn_weights, present_key_value = self.self_attn(
|
# 因为新版在 output_attentions=False 时可能只返回 (hidden_states, present_key_value)
|
||||||
|
attn_outputs = self.self_attn(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
attention_mask=attention_mask,
|
attention_mask=attention_mask,
|
||||||
position_ids=position_ids,
|
position_ids=position_ids,
|
||||||
past_key_value=past_key_value,
|
past_key_value=past_key_value,
|
||||||
output_attentions=output_attentions,
|
output_attentions=output_attentions,
|
||||||
use_cache=use_cache,
|
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
|
hidden_states = residual + hidden_states
|
||||||
|
|
||||||
# Fully Connected
|
# Fully Connected
|
||||||
@@ -129,7 +130,15 @@ class DiffLlama(LlamaModel):
|
|||||||
dropout=0.1,
|
dropout=0.1,
|
||||||
ffn_dropout=0.1,
|
ffn_dropout=0.1,
|
||||||
attention_dropout=0.0,
|
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)
|
super().__init__(config)
|
||||||
|
|
||||||
@@ -293,6 +302,13 @@ class DiffLlama(LlamaModel):
|
|||||||
else:
|
else:
|
||||||
position_ids = position_ids.view(-1, seq_length).long()
|
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
|
# embed positions
|
||||||
if attention_mask is None:
|
if attention_mask is None:
|
||||||
attention_mask = torch.ones(
|
attention_mask = torch.ones(
|
||||||
@@ -300,6 +316,9 @@ class DiffLlama(LlamaModel):
|
|||||||
dtype=torch.bool,
|
dtype=torch.bool,
|
||||||
device=inputs_embeds.device,
|
device=inputs_embeds.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 处理 attention_mask 格式
|
||||||
|
# 注意:某些新版本 transformers 此处返回值结构有变,如报错请检查此函数
|
||||||
attention_mask = self._prepare_decoder_attention_mask(
|
attention_mask = self._prepare_decoder_attention_mask(
|
||||||
attention_mask,
|
attention_mask,
|
||||||
(batch_size, seq_length),
|
(batch_size, seq_length),
|
||||||
@@ -329,23 +348,11 @@ class DiffLlama(LlamaModel):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.gradient_checkpointing and self.training:
|
if self.gradient_checkpointing and self.training:
|
||||||
|
# 原始代码此处有 NotImplementedError,如果需要开启梯度检查点,
|
||||||
|
# 也需要将 position_embeddings 传入 custom_forward
|
||||||
raise NotImplementedError
|
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:
|
else:
|
||||||
|
# [MODIFIED] 在调用 decoder_layer 时传入 position_embeddings
|
||||||
layer_outputs = decoder_layer(
|
layer_outputs = decoder_layer(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
attention_mask=attention_mask,
|
attention_mask=attention_mask,
|
||||||
@@ -354,6 +361,7 @@ class DiffLlama(LlamaModel):
|
|||||||
output_attentions=output_attentions,
|
output_attentions=output_attentions,
|
||||||
use_cache=use_cache,
|
use_cache=use_cache,
|
||||||
cond_embedding=diffusion_step,
|
cond_embedding=diffusion_step,
|
||||||
|
position_embeddings=position_embeddings, # 新增参数
|
||||||
)
|
)
|
||||||
|
|
||||||
hidden_states = layer_outputs[0]
|
hidden_states = layer_outputs[0]
|
||||||
@@ -375,14 +383,6 @@ class DiffLlama(LlamaModel):
|
|||||||
|
|
||||||
hidden_states = self.mel_out_mlp(hidden_states)
|
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:
|
if return_dict:
|
||||||
return {
|
return {
|
||||||
"output": hidden_states,
|
"output": hidden_states,
|
||||||
|
|||||||
Reference in New Issue
Block a user