You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.29 21:07:38