运行GitHub上的PyTorch WGAN代码遇AttributeError及TensorBoardX报错
Fixing Two Common Errors in improved-wgan-pytorch
1. AttributeError: 'function' object has no attribute 'Variable'
这个问题纯粹是PyTorch版本迭代搞出来的——在PyTorch 0.4.0之后,Variable和Tensor被合并了,现在根本不需要用Variable来包裹张量实现自动求导。旧代码里的from torch.autograd import Variable以及Variable(tensor)写法,在新版PyTorch里就会触发这个错误,因为torch.autograd.Variable现在是个函数,不是类属性了。
修复方法很简单:
- 找到代码里所有导入
Variable的语句,比如from torch.autograd import Variable,直接删掉就行。 - 把所有
Variable(xxx)的调用换成xxx本身。如果需要让张量支持自动求导,要么在创建时加requires_grad=True,要么调用xxx.requires_grad_(True)。
举个例子,旧代码:
改成:x = Variable(torch.randn(3, 3), requires_grad=True)x = torch.randn(3, 3, requires_grad=True)
2. Error in writer.add_scalar('data/disc_cost', disc_cost, iteration)
这个报错一般有两个常见原因,对应不同的修复方式:
原因1:disc_cost是张量而非标量值
tensorboardX的add_scalar需要的是一个普通数值(比如int或float),不是PyTorch张量。哪怕disc_cost是形状为(1,)的张量,直接传进去也会报错。
解决办法:
把disc_cost换成disc_cost.item(),取出张量里的具体数值:
writer.add_scalar('data/disc_cost', disc_cost.item(), iteration)
要是你用的是特别老的PyTorch版本(不过看第一个错误应该是新版本),也可以用disc_cost.data[0],但item()是现在官方推荐的写法。
原因2:tensorboardX版本和PyTorch不兼容
要是已经用了item()还是报错,那大概率是tensorboardX版本太旧,适配不了你当前的PyTorch版本。
解决办法:
- 先卸载旧版本:
pip uninstall tensorboardX -y - 安装最新版:
pip install tensorboardX --upgrade - 或者直接换成PyTorch官方的
tensorboard(现在官方已经内置支持,不用再依赖tensorboardX了):
然后把代码里导入tensorboardX的部分改成:pip install tensorboard
后面的from torch.utils.tensorboard import SummaryWriterwriter用法和之前完全一样,不用改其他代码。
内容的提问来源于stack exchange,提问作者Shan
相关产品推荐
相关产品推荐

