
3条命令把openpi的JAX模型搬进PyTorch【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi还在为JAX检查点进不了PyTorch生态发愁openpi 是 Physical Intelligence 开源的机器人 VLA视觉-语言-动作模型仓库它把 π₀ / π₀.₅ 的模型转换工具直接写进了 examples/convert_jax_model_to_pytorch.py一条命令即可把 JAX 检查点导出成 PyTorch 权重之后推理、微调、起服务全部无缝切换。读完本文你将拿到从环境准备到权重导出的 4 步可复制命令2 处最关键的参数重排逻辑以及为什么会这么写3 个高频坑的最短修复命令加一段最小校验代码整条链路很短脚本先用 orbaxJAX 的 checkpoint 读写库把检查点恢复成纯参数字典再分别重排视觉塔和专家层的权重最后实例化 PyTorch 版模型并落盘。 第一步装好环境并确认 transformers 版本先克隆仓库记得带子模块再用 uv 建环境。PyTorch 实现依赖打过补丁的 transformers版本必须是 4.53.2。git clone --recurse-submodules https://gitcode.com/GitHub_Trending/op/openpi cd openpi GIT_LFS_SKIP_SMUDGE1 uv sync GIT_LFS_SKIP_SMUDGE1 uv pip install -e . uv pip show transformers确认版本后把补丁文件覆盖进虚拟环境里的 transformerscp -r ./src/openpi/models_pytorch/transformers_replace/* .venv/lib/python3.11/site-packages/transformers/Name: transformers Version: 4.53.2注意uv 默认用硬链接模式这个覆盖会污染 uv 缓存。想彻底撤销时执行uv cache clean transformers。 第二步用 inspect_only 预览 JAX 参数结构正式转换前先看看检查点里有哪些参数、各是什么形状。--inspect_only只读不写输出格式为参数名: (shape)dtype来自 src/openpi/training/utils.py 的array_tree_to_info。uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --inspect_only输出节选img/embedding/kernel: (14, 14, 3, 1152)float32 llm/embedder/input_embedding: (257152, 2048)float32 ...官方检查点默认缓存在~/.cache/openpi下用环境变量OPENPI_DATA_HOME可以改位置。⚡ 第三步3个参数一键导出 PyTorch 权重核心就 2 个必填参数加 1 个输出目录。--config_name取自 src/openpi/training/config.py 里注册的模型名比如pi0_droid、pi05_droid、pi0_aloha_sim。uv run examples/convert_jax_model_to_pytorch.py \ --checkpoint_dir ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid \ --config_name pi0_droid \ --output_path ~/.cache/openpi/openpi-assets/checkpoints/pi0_droid_pytorchConverting PI0 checkpoint from .../pi0_droid to .../pi0_droid_pytorch Model config: Pi0Config(...) Model conversion completed successfully! Model saved to .../pi0_droid_pytorch输出目录里会得到 3 样东西model.safetensors全部权重、config.json记录 action_dim、action_horizon 等 5 个字段、assets/归一化统计从检查点的上一级目录自动复制。精度默认 bfloat16需要全精度时加--precision float32。✅ 第四步指向转换后目录直接起策略服务PyTorch 权重的推理入口和 JAX 版完全一样你只改检查点路径。起服务时把--policy.dir指向转换输出即可uv run scripts/serve_policy.py policy:checkpoint \ --policy.configpi0_droid \ --policy.dir/path/to/pi0_droid_pytorch服务启动后监听 8000 端口等待观测数据机器人端如何接入可看 docs/remote_inference.md。后面要用 scripts/train_pytorch.py 做 PyTorch 微调时这个转换好的目录就是基础模型权重来源。 细节解析两处最容易被忽略的参数处理为什么先过一遍 JAX 加载器检查点存储时的 dtype 和模型恢复时的 dtype 可能不一致。slice_initial_orbax_checkpoint直接调 src/openpi/models/model.py 的恢复逻辑让权重先走一遍 JAX 模型加载dtype 转换和真实训练时完全一致# examples/convert_jax_model_to_pytorch.py - slice_initial_orbax_checkpoint params openpi.models.model.restore_params( f{checkpoint_dir}/params/, restore_typenp.ndarray, dtyperestore_precision )pi05 的归一化为什么单独分支pi0 的层归一化是只有 scale 的 RMSNorm而 pi05 换成了带 kernel bias 的自适应 Dense 层。slice_gemma_state_dict用检查点路径里是否含pi05字符串来区分两套参数名# examples/convert_jax_model_to_pytorch.py - slice_gemma_state_dict if pi05 in checkpoint_dir: llm_input_layernorm_bias state_dict.pop(fllm/layers/pre_attention_norm_{num_expert}/Dense_0/bias{suffix}) else: llm_input_layernorm state_dict.pop(fllm/layers/pre_attention_norm_{num_expert}/scale{suffix})视觉塔这边patch 嵌入的卷积核按transpose(3, 2, 0, 1)从 JAX 的 (H, W, C_in, C_out) 转成 PyTorch 的 (C_out, C_in, H, W)这也是两个框架最常见的维度差异来源。转换前必看的3个坑坑 1config_name 传了 fast 系列现象ValueError: Config pi0_fast_droid is not a Pi0Config原因脚本只接受 π₀ / π₀.₅ 的 Pi0ConfigPyTorch 版暂不支持 π₀-FAST。修复换成pi0_droid、pi05_droid等 pi0/pi05 配置名。坑 2重命名了检查点目录现象pi05 检查点转换后归一化层权重对不上推理行为异常。原因脚本靠pi05 in checkpoint_dir判断走哪套归一化分支目录名去掉 pi05 字样就会走错分支。修复目录名保留 pi05如~/.cache/openpi/openpi-assets/checkpoints/pi05_droid。坑 3输出目录缺 assets/现象转换本身成功但推理加载策略时报缺归一化统计的错误。原因脚本只从检查点上一级目录找assets/找不到就静默跳过归一化统计没跟过来。修复转换前把源检查点的assets/放到checkpoint_dir的父目录或转换后手动复制进输出目录。最小校验确认 PyTorch 权重真的能被加载在仓库根目录跑这段代码src/openpi/policies/policy_config.py 的create_trained_policy会通过目录里有没有model.safetensors自动切换 PyTorch 加载路径from openpi.training import config as _config from openpi.policies import policy_config config _config.get_config(pi0_droid) policy policy_config.create_trained_policy(config, /path/to/pi0_droid_pytorch)怎么算通过日志出现Loading model...且没有缺键报错说明权重被 src/openpi/models_pytorch/pi0_pytorch.py 的PI0Pytorch完整接收再喂一组观测跑policy.infer(example)[actions]输出第一步动作序列长度应为 10pi0_droid的 action_horizonpi05_droid是 15数值落在正常关节角度范围内即视为转换成功。收个尾转换全程 3 个参数--checkpoint_dir、--config_name、--output_path精度默认 bfloat16目录名别乱改pi05 的分支判断依赖路径字符串assets 也藏在检查点父目录输出目录即插即用推理、起服务、PyTorch 微调都指向它下一篇我们拆train_pytorch.py的单卡与多卡微调流程。有问题可以按 CONTRIBUTING.md 的指引提交 issue 或贡献补丁。【免费下载链接】openpi项目地址: https://gitcode.com/GitHub_Trending/op/openpi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考