为何PyTorch BatchNorm1D对整数型张量归一化时报Long类型未实现错误?
PyTorch BatchNorm1d 整数张量报错解决
问题现象
使用PyTorch的nn.BatchNorm1d对整数型张量做归一化时触发RuntimeError,错误提示"batch_norm" not implemented for 'Long',但将张量改为浮点型时可正常运行。
错误复现代码
import torch import torch.nn as nn # 整数型张量 test_int_input = torch.randint(size = [3,5],low=1,high=9) # BatchNorm1D 实例 batchnorm1D = nn.BatchNorm1d(num_features=5) test_output = batchnorm1D(test_int_input)
报错信息
--------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) <ipython-input-38-6c672cd731fa> in <module> 1 batchnorm1D = nn.BatchNorm1d(num_features=5) ----> 2 test_output = batchnorm1D(test_input) /opt/conda/lib/python3.7/site-packages/torch/nn/modules/module.py in __call__(self, *input, **kwargs) 530 result = self._slow_forward(*input, **kwargs) 531 else: ---> 532 result = self.forward(*input, **kwargs) 533 for hook in self._forward_hooks.values(): 534 hook_result = hook(self, input, result) /opt/conda/lib/python3.7/site-packages/torch/nn/modules/batchnorm.py in forward(self, input) 105 input, self.running_mean, self.running_var, self.weight, self.bias, 106 self.training or not self.track_running_stats, ---> 107 exponential_average_factor, self.eps) 108 109 /opt/conda/lib/python3.7/site-packages/torch/nn/functional.py in batch_norm(input, running_mean, running_var, weight, bias, training, momentum, eps) 1668 return torch.batch_norm( 1669 input, weight, bias, running_mean, running_var, -> 1670 training, momentum, eps, torch.backends.cudnn.enabled 1671 ) 1672 RuntimeError: "batch_norm" not implemented for 'Long'
正常运行示例(浮点型张量)
import torch import torch.nn as nn # 浮点型张量 test_input = torch.randn(size = [3,5]) # BatchNorm1D 实例 batchnorm1D = nn.BatchNorm1d(num_features=5) test_output = batchnorm1D(test_input) test_output
输出结果
tensor([[ 0.4311, -1.1987, 0.9059, 1.1424, 1.2174], [-1.3820, 1.2492, -1.3934, 0.1508, 0.0146], [ 0.9509, -0.0505, 0.4875, -1.2931, -1.2320]], grad_fn=<NativeBatchNormBackward>)
解决方法
BatchNorm1d的计算依赖均值、方差的统计,涉及除法、平方根等浮点运算,而整数类型张量(LongTensor)不支持这些操作。只需将整数张量转换为浮点类型即可解决:
import torch import torch.nn as nn # 将整数张量转为浮点型(使用.float()方法) test_int_input = torch.randint(size=[3,5], low=1, high=9).float() batchnorm1D = nn.BatchNorm1d(num_features=5) test_output = batchnorm1D(test_int_input) print(test_output)
也可以指定具体的浮点类型,比如:
test_int_input = torch.randint(size=[3,5], low=1, high=9).to(torch.float32)
内容的提问来源于stack exchange,提问作者Aditya Kansal
相关产品推荐
相关产品推荐

