PyTorch报错:torch.set_grad_enabled提示'bool'对象不可调用
PyTorch
torch.set_grad_enabled(False) 报错:TypeError: 'bool' object is not callable 解决方法 核心问题是你在代码中直接覆盖了PyTorch的torch.set_grad_enabled函数,将其赋值为布尔值,导致后续调用时触发类型错误。以下是具体错误点和修复方案:
1. 致命错误:覆盖torch.set_grad_enabled函数
在CardAgent的train_memory方法中,你写了:
torch.set_grad_enabled = True
这行代码把原本是函数的torch.set_grad_enabled直接改成了布尔值True。第一次运行时还能正常执行,但第二次调用torch.set_grad_enabled(False)时,它已经变成了布尔值,自然无法被调用。
修复方法:将赋值改为函数调用:
torch.set_grad_enabled(True)
或者更规范地使用上下文管理器(和你在play函数里的with torch.no_grad()逻辑统一):
with torch.set_grad_enabled(True): # 所有训练相关代码放在这里
2. 其他潜在错误修复
(1) 错误设置模型梯度开关
在network方法中:
self.requires_grad_ = False
这同样是把模型的requires_grad_方法覆盖成了布尔值,正确写法是调用方法:
self.requires_grad_(False)
注:如果只是推理时需要禁用梯度,建议用torch.no_grad()上下文管理器,而非直接冻结整个模型参数(会影响后续训练)。
(2) Forward方法中的重复计算
在forward方法里,你连续两次调用self.fc4(x),导致第一次的ReLU结果被丢弃:
x = F.relu(self.fc4(x)) x = F.softmax(self.fc4(x), dim=-1)
修复方法:复用计算结果:
x = F.relu(self.fc4(x)) x = F.softmax(x, dim=-1)
(3) 训练时的张量错误
在train_memory中,你错误地将next_state替换成了observation:
next_state_tensor = torch.tensor(np.expand_dims(observation, 0), dtype=torch.float32, requires_grad = True)
修复方法:改为使用next_state,且无需手动设置requires_grad=True(模型梯度开关由torch.set_grad_enabled控制):
next_state_tensor = torch.tensor(np.expand_dims(next_state, 0), dtype=torch.float32)
修复后的关键代码片段
修复后的train_memory方法
def train_memory(self, observation, move, reward, next_state, complete): self.train() # 修复:调用函数而非赋值 torch.set_grad_enabled(True) state_tensor = torch.tensor(np.expand_dims(observation, 0), dtype=torch.float32) # 修复:使用正确的next_state next_state_tensor = torch.tensor(np.expand_dims(next_state, 0), dtype=torch.float32) if not complete: target = reward + self.gamma * torch.max(self.forward(next_state_tensor[0])) output = self.forward(state_tensor) target_f = output.clone() target_f[0][np.argmax(move)] = target target_f.detach() self.optimizer.zero_grad() loss = F.mse_loss(output, target_f) loss.backward() self.optimizer.step()
修复后的forward方法
def forward(self, observation): x = F.relu(self.fc1(observation)) x = F.relu(self.fc2(x)) x = F.relu(self.fc3(x)) x = F.relu(self.fc4(x)) # 修复:复用fc4的输出结果 x = F.softmax(x, dim=-1) if self.mask is not None: print(x) return x * self.mask return x
内容的提问来源于stack exchange,提问作者aimkeys mwaura
相关产品推荐
相关产品推荐

