PyTorch加载自定义Dataset时出现NumPy数组不可写警告的原因及解决方法
告警原因
UserWarning: The given NumPy array is not writeable, and PyTorch does not support non-writeable tensors.
这个告警的触发逻辑如下:
mpimg.imread()加载返回的图像是numpy数组,后续调用np.transpose()做轴重排时,该操作返回的是原数组的视图而非独立副本,这类numpy视图默认被标记为非可写状态。- PyTorch的
DataLoader默认default_collate函数在把numpy数组转为张量时,会优先共享内存避免拷贝,但PyTorch本身不支持非可写张量,两者冲突触发告警。该告警的潜在风险是:后续对生成的PyTorch张量做写入操作时,会直接修改原本标记为只读的numpy底层数组,可能引发不可预期的异常。
修复方法
以下三种方案任选其一即可,优先推荐第一种,逻辑最简单无副作用:
- 方案1:转置后生成numpy数组的可写副本,仅需修改
np.transpose行代码:
直接得到独立的可写numpy数组,转张量时不会触发告警,小尺寸图像的拷贝开销可以忽略。positive_img = np.transpose(positive_img, (2,0,1)).copy() - 方案2:返回前主动将numpy数组转为PyTorch张量,由PyTorch处理内存逻辑:
positive_img = torch.from_numpy(positive_img) return positive_img - 方案3:手动修改numpy数组的可写标记,适合不想额外拷贝内存的场景:
注意该方案下如果原数组有其他引用,修改内容会同步生效。positive_img = np.transpose(positive_img, (2,0,1)) positive_img.flags.writeable = True
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

