pytorch-image-models
在昇腾 NPU 上安装 timm,并用预训练 ResNet-18 跑通图像分类推理和多尺度特征提取。
前置条件
硬件
Atlas 900 A2 / A3 或 Ascend 950 系列服务器(Ascend 910B),并按需完成物理机或容器内的设备挂载。
基础软件
在运行本文档示例之前,你的机器上需要已经装好并可用:
可用的 Python 环境
可用的 CANN(参考快速安装昇腾环境)
本文档示例在 Python 3.12、CANN 9.1.0 环境下验证通过。
本文档配套镜像:swr.cn-south-1.myhuaweicloud.com/ascendhub/cann:9.1.0-910b-ubuntu22.04-py3.12。
加载 CANN 环境
source /usr/local/Ascend/ascend-toolkit/set_env.sh
安装 PyTorch 软件栈
安装与 CANN 配套的 PyTorch、Torch-NPU 和 TorchVision,并查看安装版本:
pip install torch==2.9.0 torch_npu==2.9.0.post2
pip install torchvision==0.24.0
python -c "import torch, torch_npu, torchvision; print('torch', torch.__version__); print('torch_npu', torch_npu.__version__); print('torchvision', torchvision.__version__)"
输出结果如下:
...
torch 2.9.0+cpu
torch_npu 2.9.0.post2
torchvision 0.24.0
安装 timm
从 PyPI 安装最新 release,并确认可以正常导入:
pip install timm
python -c "import timm; print('timm', timm.__version__)"
输出结果如下,其中 xxx 表示实际版本号:
...
timm xxx
快速开始
示例 1:图像分类推理
timm 的入门示例是 timm.create_model('resnet18') 创建模型并前向推理。详见官方文档。
首次运行自动从 ModelScope 下载 ResNet-18 预训练权重(来源 timm/resnet18.a1_in1k),输入随机图像,输出 1000 类 logits,运行下面的 Python 脚本:
import os
import safetensors.torch
import timm
import torch
import torch_npu
try:
from modelscope.hub.snapshot_download import snapshot_download
model_dir = snapshot_download("timm/resnet18.a1_in1k")
model = timm.create_model("resnet18")
model.load_state_dict(
safetensors.torch.load_file(os.path.join(model_dir, "model.safetensors"))
)
except Exception:
os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
model = timm.create_model("resnet18", pretrained=True)
model = model.to("npu:0").eval()
x = torch.randn(1, 3, 224, 224, device="npu:0")
with torch.no_grad():
out = model(x)
print("out_shape", tuple(out.shape))
输出结果如下:
out_shape (1, 1000)
示例 2:多尺度特征提取
timm 的 features_only=True 可将任意模型转为多尺度特征提取器,输出各层 feature map。详见Feature Extraction 文档。
加载预训练 ResNet-18,输出 stride 2/4/8/16/32 共 5 层 feature map,打印每层形状和通道数,运行下面的 Python 脚本:
import os
import safetensors.torch
import timm
import torch
import torch_npu
try:
from modelscope.hub.snapshot_download import snapshot_download
model_dir = snapshot_download("timm/resnet18.a1_in1k")
model = timm.create_model("resnet18", features_only=True)
model.load_state_dict(
safetensors.torch.load_file(os.path.join(model_dir, "model.safetensors")),
strict=False,
)
except Exception:
os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
model = timm.create_model("resnet18", features_only=True, pretrained=True)
model = model.to("npu:0").eval()
x = torch.randn(1, 3, 224, 224, device="npu:0")
with torch.no_grad():
outs = model(x)
print("num_features", len(outs))
print("channels", model.feature_info.channels())
for i, o in enumerate(outs):
print(i, tuple(o.shape))
输出结果如下:
num_features 5
channels [64, 64, 128, 256, 512]
0 (1, 64, 112, 112)
1 (1, 64, 56, 56)
2 (1, 128, 28, 28)
3 (1, 256, 14, 14)
4 (1, 512, 7, 7)