z
zyzoe/ppo-Pyramids
模型介绍
文件和版本
Pull Requests
讨论
分析

ppo-Pyramids (Nablaaa/ppo-Pyramids) 昇腾 NPU 适配

Unity ML-Agents PPO 训练的 Pyramids 金字塔导航策略(2×512 MLP,172 维向量观测:raycast 感知 + 目标 one-hot),本仓将其适配到华为昇腾 NPU 真实跑通推理。

验证环境

  • NPU: Ascend910 (npu:0),CANN 8.5.1,npu-smi 25.5.5
  • 推理引擎: torch_npu(torch 2.1.0 + torch_npu 2.1.0.post3),fp32,无 CPU fallback
  • 依赖: 仅 torch + torch_npu(纯 nn.Module 同构重建,无第三方 RL 库)
  • 权重: /opt/atomgit/.hfhome/ppo-Pyramids/agent.pt(ML-Agents Pyramids/checkpoint.pt,global_step 3,778,906,8.7 MB,权重不入库)

网络结构(从权重读出)

  • 观测维度: 172(network_body._body_endoder.seq_layers.0.weight 形状 (512, 172))
  • 编码器: Linear(172→512) → swish → Linear(512→512)(config.json hidden_units=512, num_layers=2, normalize=false, vis_encode_type=simple 无视觉编码器)
  • 动作输出: 单离散分支 5 动作(action_model._discrete_distribution.branches.0.weight 形状 (5, 512))

权重加载门禁(真机实测)

  • gate1: 同构 PyramidsPolicy(含 ML-Agents 元信息 buffer)+ load_state_dict(strict=True) → missing=0 / unexpected=0 / mismatched=0
  • 参数量: 353,797(88,576 + 262,656 + 2,565 = 编码器 2×512 + 5 动作离散头)

真机推理结果(npu:0)

合成 8 组 172 维 Pyramids 观测(linspace(-1,1) 覆盖 raycast 距离/方向取值域,含全零基线):

指标值
动作输出[2, 4, 4, 4, 1, 1, 1, 1](不同感知输入 → 不同导航动作)
动作多样性3 种 / 8 组观测(非恒定,策略真实响应观测)
top-1 置信度0.863(softmax 平均)

性能(真机实测)

输入p50avg
1×172 观测, fp320.148 ms0.150 ms

运行方式

export PYTHONPATH=/tmp/fixpkgs            # 先于 set_env.sh
source /usr/local/Ascend/cann-8.5.1/set_env.sh
export ASCEND_RT_VISIBLE_DEVICES=0
python inference.py

测试用例

  1. strict load 门禁:同构网络严格加载 checkpoint,键/形状零差异
  2. 动作非恒定:8 组覆盖观测域的输入产生 ≥2 种动作(实际 3 种)
  3. 语义验证:负值观测偏动作 4/正值偏动作 1,随感知输入单调变化
  4. 性能:batch=1 前向 p50 < 1 ms