PyTorch多Actor头PPO实现梯度计算原地操作错误排查
统一张量类型,避免混合精度冲突
报错里的torch.DoubleTensor版本不符,先查模型参数和输入obs的 dtype 是否一致。跑print(next(model.parameters()).dtype)和print(obs.dtype),把所有张量统一成torch.float32或者torch.float64——混合类型会触发隐式转换,很容易触发原地修改的检测报错。检查共享层的操作逻辑
多Actor头结构里,共享层的输出会被多个头部复用,绝对不能对共享层输出做原地修改操作,比如shared_out += x这种,必须改成shared_out = shared_out + x。另外PPO里新旧策略的共享层参数要彻底分开,不能直接复用旧策略的张量,用copy.deepcopy复制参数或者重新初始化,不然反向传播时会不小心修改旧策略的张量。暂时关闭自动混合精度
PyTorch 2.2.2的自动混合精度在多分支结构下可能出问题,如果你开了torch.cuda.amp.autocast(),先关掉试试,看错误会不会消失。动作分支的切片操作要克隆张量
动作选择头输出的索引用来切共享层输出时,别直接用selected_out = shared_out[:, action_idx],要先克隆:selected_out = shared_out[:, action_idx].clone()——直接切片得到的张量是原张量的视图,后续修改会触发原地修改的报错。用详细梯度检测定位报错行
除了默认的异常检测,把前向和反向传播的代码用torch.autograd.detect_anomaly(True)包裹,能精准定位到哪一行触发了原地修改:with torch.autograd.detect_anomaly(True): # 前向传播算动作和log概率 action_type, action_params = model(obs) # 计算PPO损失 loss = compute_ppo_loss(action_type, action_params, old_log_probs, returns) # 反向传播 loss.backward()
内容的提问来源于stack exchange,提问作者Xandra Dave Cochran

