PyTorch+NumPy训练代码报错:'numpy.ndarray'对象没有'unsqueeze'属性
解决
AttributeError: 'numpy.ndarray' object has no attribute 'unsqueeze'错误 这个问题很直观——你正在尝试对numpy数组调用只有PyTorch张量才支持的unsqueeze()方法,工具不对自然会触发报错。
问题根源
你的代码里config1.both_transform(image=np.asarray(image))["image"]返回的结果是numpy.ndarray类型,而unsqueeze()是PyTorch torch.Tensor类的专属方法,numpy数组根本没有这个属性,所以才会抛出这个错误。
修复方案
只需要把numpy数组转换成PyTorch张量,再执行后续操作即可,具体修改后的代码如下:
def plot_example(low_res_folder, gen): files=os.listdir(low_res_folder) gen.eval() for file in files: image=Image.open("test_images/" + file) with torch.no_grad(): # 1. 获取变换后的numpy数组 transformed_image_np = config1.both_transform(image=np.asarray(image))["image"] # 2. 转换为PyTorch张量,同时调整数据类型为float(模型通常用float32) transformed_image_tensor = torch.from_numpy(transformed_image_np).float() # 3. 增加batch维度并移动到指定设备 upscaled_img=gen( transformed_image_tensor.unsqueeze(0).to(config1.DEVICE) ) save_image(upscaled_img * 0.5 + 0.5, f"saved/{file}") gen.train()
如果你喜欢更紧凑的写法,也可以把转换步骤合并成一行:
upscaled_img=gen( torch.from_numpy(config1.both_transform(image=np.asarray(image))["image"]) .float() .unsqueeze(0) .to(config1.DEVICE) )
额外提示
如果你的both_transform输出的是0-255范围的uint8类型数组,转成float后建议做归一化处理(比如除以255),确保输入数据的范围和你的生成器模型训练时的输入范围一致,避免影响最终的生成效果。
内容的提问来源于stack exchange,提问作者Khalid El amraoui
相关产品推荐
相关产品推荐

