ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

如何把 PaddleClas 的 PP-LCNetV2 权重转换成 timm 格式并在 create_model 中加载?

如何把 PaddleClas 的 PP-LCNetV2 权重转换成 timm 格式并在 create_model 中加载? 如何把 PaddleClas 的 PP-LCNetV2 权重转换成 timm 格式并在 create_model 中加载【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models手头有一个 PaddleClas 官方发布的 PP-LCNetV2 检查点.pdparams文件想把它用在 PyTorch 生态里例如加载进 timm 的模型做推理或微调。pytorch-image-modelstimm仓库提供了专门的转换脚本 convert/convert_lcnetv2_paddle.py可以把 PaddleClas 的权重转成 timm 可直接加载的 state dict 文件再用timm.create_model的checkpoint_path参数加载。本文按准备 → 转换 → 验证 → 加载这条路径走一遍。准备条件安装 paddlepaddle 与 timm转换脚本本身是纯 Python PyTorch但读取.pdparams文件需要paddlepaddle。脚本头部明确注明paddlepaddleis required to unpickle the .pdparams files, it is not in requirements.txt也就是说paddlepaddle不在项目的 requirements.txt 里需要单独安装。timm 本体按项目常规的pip install timm或从本仓库安装即可。转换前还需要确认两件事你拿到的是 PaddleClas 官方 PP-LCNetV2 检查点.pdparams格式对应的 timm 模型名。timm 中注册的 PP-LCNetV2 入口有 timm/models/lcnetv2.py 里的lcnetv2_small、lcnetv2_base、lcnetv2_large三个--model参数默认就是lcnetv2_base要和你的权重规模对上。执行转换在仓库根目录下运行转换脚本python convert/convert_lcnetv2_paddle.py PPLCNetV2_base_pretrained.pdparams \ --model lcnetv2_base --output lcnetv2_base.pth其中PPLCNetV2_base_pretrained.pdparams是脚本文档中给出的示例文件名替换成你实际下载的 PaddleClas 检查点路径--output指定转换产物路径不传时默认写到./converted.pth。参数说明checkpoint位置参数PaddleClas 的.pdparams检查点路径--model目标 timm 模型名默认lcnetv2_base--output输出文件路径默认./converted.pth--dropout-probPaddle 侧建模时使用的 dropout 概率默认0.2。脚本帮助文本提示该值要与 PaddleClas 的PPLCNetV2_*入口点保持一致——Paddle 在推理时按1 - p缩放 pre-logits 特征转换脚本把这个缩放因子直接乘进last_conv.weight里如果这里传的 p 和原模型不一致logits 会整体偏小。转换过程中脚本会做一次内置校验用timm.create_model(args.model, pretrainedFalse)构建目标模型然后对转换出的 state dict 执行model.load_state_dict(state_dict)严格模式key 不匹配会直接抛异常通过后才写出文件。也就是说只要脚本正常跑完key 与形状就已经和 timm 模型对上了。验证转换结果脚本成功时的输出格式为Converted {len(state_dict)} tensors from {args.checkpoint} to {args.output}即打印实际转换的张量数量、源文件与目标文件路径上方为文档示例格式具体张量数以实际运行为准。同时确认输出目录里生成了--output指定的.pth文件。在 create_model 中加载转换后的权重timm 的工厂函数 timm/models/_factory.py 中create_model支持checkpoint_path参数其文档说明是Path of checkpoint to loadafterthe model is initialized模型初始化完成后调用load_checkpoint(model, checkpoint_path)载入权重。所以加载方式就是import timm model timm.create_model(lcnetv2_base, checkpoint_pathlcnetv2_base.pth)注意这里不再传pretrainedTrue那会拉取 ImageNet-1k 预训练权重转换得到的 checkpoint 通过checkpoint_path加载。lcnetv2_*系列默认 1000 类、输入 224×224、crop_pct 0.875见 timm/models/lcnetv2.py 的_cfg如需其他类别数可传num_classes但此时 classifier 权重不会被 checkpoint 覆盖。仓库根目录的 inference.py 也走同一条加载路径--checkpoint参数传入检查点路径--model指定架构例如python inference.py --data-dir /path/to/imagenet/val --model lcnetv2_base --checkpoint lcnetv2_base.pth该脚本在指定了--checkpoint时不会使用预训练权重args.pretrained args.pretrained or not args.checkpoint输出 top-k 类别结果到控制台和 CSV。可选推理前做重参数化PP-LCNetV2 的 depthwise 卷积在训练时是可重参数化的多尺度分支如 5×5、3×3 并行分支timm 侧保留了这些分支结构官方发布的 Paddle 检查点里同时带有折叠后的单一 depthwise 卷积转换脚本会跳过它只保留分支。按脚本注释timm 的做法是按需折叠如果要在 CPU 上做推理、获得折叠后的单一卷积结构可以用from timm.utils import reparameterize_model model reparameterize_model(model)timm/utils/model.py 中的reparameterize_model会遍历模型对实现了reparameterize的子模块PP-LCNetV2 的RepDepthwiseSeparable即属此类调用其reparameterize()把多分支折叠进单个 depthwise 卷积。不折叠也能正常前向只是推理时走分支求和路径。适用范围与限制转换脚本只针对 PaddleClas 的 PP-LCNetV2 检查点其他 Paddle 权重不能直接套用脚本假设发布版检查点同时携带重参数化 dw_conv 与分支权重见脚本内注释非官方发布的训练中间态检查点未经验证--dropout-prob必须与 Paddle 侧建模型时的取值一致脚本默认 0.2 对应官方 PPLCNetV2 入口点的默认设置若你用的是自己调过该参数的模型需要显式传入对应值。完成转换并通过create_model加载后模型即可按 timm 常规方式参与推理、微调或导出仓库的onnx_export.py等脚本同样接受--model lcnetv2_base加--checkpoint的工作方式。【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表