You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

微调Gemma2-2B时添加attention_mask后Loss出现NaN问题求助

Gemma2-2B序列分类任务传入attention_mask后Loss出现NaN的解决方法

问题背景

使用Gemma2-2B结合LoRA微调序列分类任务时,当输入做了max_length padding并传入attention_mask,且模型使用bfloat16精度时,模型输出的Loss和logits全部变为NaN。但以下场景下结果正常:

  • 不传入attention_mask
  • 输入不做padding(attention_mask全为1)
  • 模型改用float16精度

可能原因

  1. bfloat16精度限制:bfloat16仅保留7位有效数字,在处理padding后的mask计算时,低精度导致的数值误差被LoRA模块的参数更新放大,触发NaN。
  2. 分类头池化逻辑:Gemma2默认用<bos> token的隐藏状态做分类,当存在padding时,mask未被正确应用于池化过程,无效的隐藏状态参与计算引发数值不稳定。
  3. LoRA配置参数:过高的LoRA秩(r)或alpha值可能导致梯度波动过大,在低精度下更容易溢出。

解决方案

1. 混合精度训练平衡精度与稳定性

如果必须使用bfloat16,可开启PyTorch自动混合精度(AMP),通过梯度缩放避免数值溢出:

from torch.cuda.amp import GradScaler, autocast

# 初始化优化器和梯度缩放器
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5)
scaler = GradScaler()

# 训练循环示例
model.train()
for batch in train_dataloader:
    optimizer.zero_grad()
    # 自动混合精度上下文
    with autocast(dtype=torch.bfloat16):
        outputs = model(
            input_ids=batch["input_ids"],
            attention_mask=batch["attention_mask"],
            labels=batch["labels"]
        )
        loss = outputs.loss
    # 缩放梯度并反向传播
    scaler.scale(loss).backward()
    # 梯度裁剪进一步稳定训练
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
    scaler.step(optimizer)
    scaler.update()

2. 替换分类头为带mask的平均池化

修改分类头的池化逻辑,基于attention_mask计算有效token的平均隐藏状态,替代默认的<bos> token单token池化:

import torch
import torch.nn as nn
from transformers.models.gemma.modeling_gemma import GemmaClassificationHead

class MaskedAvgClassificationHead(GemmaClassificationHead):
    def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor = None) -> torch.Tensor:
        if attention_mask is not None:
            # 扩展mask维度以匹配隐藏状态
            mask = attention_mask.unsqueeze(-1).expand(hidden_states.shape)
            # 计算有效token的加权平均
            sum_hidden = torch.sum(hidden_states * mask, dim=1)
            sum_mask = torch.clamp(mask.sum(dim=1), min=1e-9)  # 避免除以0
            pooled_output = sum_hidden / sum_mask
        else:
            pooled_output = hidden_states[:, 0, :]  # 默认取<bos>
        # 沿用原分类头的全连接层和dropout
        pooled_output = self.dropout(pooled_output)
        logits = self.out_proj(pooled_output)
        return logits

# 替换模型的分类头
model.classifier = MaskedAvgClassificationHead(
    hidden_size=model.config.hidden_size,
    num_labels=model.config.num_labels,
    dropout=model.config.classifier_dropout
).to(model.device)

3. 调整LoRA配置降低数值波动

降低LoRA的秩和alpha值,减少参数更新带来的梯度波动:

peft_config = LoraConfig(
    task_type=TaskType.SEQ_CLS,
    inference_mode=False,
    r=4,  # 从8降低到4
    lora_alpha=16,  # 与r成比例调整
    lora_dropout=0.05,  # 降低dropout比例
    target_modules=['down_proj','o_proj','k_proj','q_proj','gate_proj','up_proj','v_proj'],
    init_lora_weights="gaussian"  # 使用高斯初始化稳定参数
)
model = get_peft_model(temp, peft_config)

4. 手动过滤padding部分梯度

在反向传播后,手动对梯度进行裁剪或过滤,避免无效梯度影响参数更新:

outputs = model(**batch)
loss = outputs.loss
loss.backward()

# 梯度裁剪限制最大范数
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

# 可选:仅保留LoRA参数的梯度,其他参数冻结
for name, param in model.named_parameters():
    if "lora" not in name:
        param.grad = None

optimizer.step()
optimizer.zero_grad()

内容的提问来源于stack exchange,提问作者Juble

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 00:20:09