torchvision transforms下nn.Sequential与Compose预处理结果不一致问题问询
问题原因
- 两条链路的插值运算的输入/输出数据类型不一致:
tr链路:Resize、CenterCrop直接操作uint8类型的torch张量,torchvision的张量插值逻辑在处理uint8输入时,输出仍为uint8,其插值计算、取整、钳位逻辑和PIL底层的BICUBIC实现存在明显差异,这部分差异在后续转float后会被Normalize放大,最终导致误差极大。tr2链路:Resize、CenterCrop操作uint8类型的PIL图像,PIL的BICUBIC插值内部会先转int32计算避免溢出,结果会先钳位到0~255再四舍五入转回uint8,和torch张量插值的uint8输出结果不一致。
- 单独测试类型转换接口无误差的原因:没有经过插值环节的uint8数据,
ToTensor和ConvertImageDtype都是直接做uint8到float32的缩放(除以255),所以结果完全一致,误差只有经过插值后的uint8数值差异才会体现。
解决方案
把ConvertImageDtype操作挪到Resize之前,让两条链路的插值都在float32的[0,1]值域上完成,规避uint8下不同插值实现的差异,修改后的tr代码如下:
tr = torch.nn.Sequential( # 先转float再做插值,和tr2的逻辑对齐 T.ConvertImageDtype(torch.float), T.Resize(224, interpolation=T.InterpolationMode.BICUBIC), T.CenterCrop(224), T.Normalize([0.48145466, 0.4578275, 0.40821073], [0.26862954, 0.26130258, 0.27577711]), )
修改后两条链路的输出平方误差和会降到1e-4以内,完全满足业务精度要求,同时保留了nn.Sequential链路无需转PIL、速度更快的优势。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

