微调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精度
可能原因
- bfloat16精度限制:bfloat16仅保留7位有效数字,在处理padding后的mask计算时,低精度导致的数值误差被LoRA模块的参数更新放大,触发NaN。
- 分类头池化逻辑:Gemma2默认用
<bos>token的隐藏状态做分类,当存在padding时,mask未被正确应用于池化过程,无效的隐藏状态参与计算引发数值不稳定。 - 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
相关产品推荐
相关产品推荐

