PyTorch Lightning训练遇RuntimeError:张量元素无梯度及grad_fn
PyTorch Lightning训练文本分类模型时触发RuntimeError问题
错误信息
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
背景
进行文本分类任务,已实现自定义数据集类UCC_Dataset、LightningDataModule子类UCC_Data_Module,模型UCC_Comment_Classifier继承自pl.LightningModule,调用trainer.fit(model, ucc_data_module)时触发上述错误。已验证所有模型参数的requires_grad均设为True,仍未找到原因。
相关代码片段
class UCC_Comment_Classifier(pl.LightningModule): def __init__(self): super().__init__() self.config = config self.pretrained_model = AutoModel.from_pretrained(config['model_name'], return_dict=True) self.hidden = torch.nn.Linear(self.pretrained_model.config.hidden_size, self.pretrained_model.config.hidden_size) self.classifier = torch.nn.Linear(self.pretrained_model.config.hidden_size, self.config['n_labels']) torch.nn.init.xavier_uniform_(self.classifier.weight) self.loss_func = nn.BCEWithLogitsLoss(reduction='mean') self.dropout = nn.Dropout() # 开启所有模型参数的梯度 for param in self.parameters(): param.requires_grad = True def forward(self, input_ids, attention_mask, labels=None): # RoBERTa层 output = self.pretrained_model(input_ids=input_ids, attention_mask=attention_mask) pooled_output = torch.mean(output.last_hidden_state, 1) # 最终logits计算 pooled_output = self.dropout(pooled_output) pooled_output = self.hidden(pooled_output) pooled_output = F.relu(pooled_output) pooled_output = self.dropout(pooled_output) logits = self.classifier(pooled_output) # 计算损失 loss = None if labels is not None: loss = self.loss_func(logits.view(-1, self.config['n_labels']), labels.view(-1, self.config['n_labels'])) return loss, logits def training_step(self, batch, batch_index): loss, outputs = self(**batch) self.log("train loss ", loss, prog_bar=True, logger=True) return {"loss": loss, "predictions": outputs, "labels": batch["labels"]} def validation_step(self, batch, batch_index): loss, outputs = self(**batch) self.log("validation loss ", loss, prog_bar=True, logger=True) return {"val_loss": loss, "predictions": outputs, "labels": batch["labels"]} def predict_step(self, batch, batch_index, dataloader_idx: int = None): loss, outputs = self(**batch) return outputs def configure_optimizers(self): optimizer = AdamW(self.parameters(), lr=self.config['lr'], weight_decay=self.config['weight_decay'], no_deprecation_warning=True) total_steps = config['train_size'] * self.config['n_epochs'] warmup_steps = math.floor(total_steps * self.config['warmup']) scheduler = get_cosine_schedule_with_warmup(optimizer, warmup_steps, total_steps) return [optimizer], [scheduler]
排查方向及解决方法
- 检查标签数据类型:BCEWithLogitsLoss要求输入的标签是浮点型张量,如果你的
labels是int型,即使能计算损失,也可能导致梯度链路异常。在数据集类中确保labels转换为torch.float32。 - 修改training_step返回方式:PyTorch Lightning更推荐直接返回loss张量而非字典,避免字典中loss的梯度被意外处理。修改为:
def training_step(self, batch, batch_index): loss, outputs = self(**batch) self.log("train loss ", loss, prog_bar=True, logger=True) return loss - 验证预训练模型梯度状态:虽然你遍历设置了
requires_grad=True,但可以在__init__末尾添加打印确认:
确保预训练模型的参数确实参与梯度更新。print("预训练模型参数是否开启梯度:", any(p.requires_grad for p in self.pretrained_model.parameters())) - 检查Loss计算的维度匹配:确认
logits和labels经过view后的形状完全一致,避免因维度不匹配导致的无梯度张量生成。 - 确认优化器覆盖所有参数:打印优化器的参数组,检查是否包含预训练模型、hidden层和classifier层的所有参数:
optimizer = AdamW(self.parameters(), lr=self.config['lr'], weight_decay=self.config['weight_decay'], no_deprecation_warning=True) print("优化器参数数量:", len(optimizer.param_groups[0]['params']))
内容的提问来源于stack exchange,提问作者hamzaouni
相关产品推荐
相关产品推荐

