本仓库对 Nablaaa/ppo-Pyramids
(Unity ML-Agents PPO 策略,Pyramids 3D 视觉环境)做昇腾(Ascend) NPU 适配,
提供纯 NPU 推理(CLI + FastAPI 双模式)。不依赖 mlagents / gym,
直接从训练检查点 .pt 手动重建策略网络,全代码零 CUDA。
Pyramids 是 Unity ML-Agents 的经典 3D 环境。本仓库导出的最终 brain 结构如下
(已与官方导出 Pyramids.onnx 用 onnxruntime CPU 逐元素核对,logits diff=0):
| 项 | 值 |
|---|---|
| 观测 obs | 172 维 float 向量(导出图内为 obs_0[56] + obs_1[56] + obs_2[56] + obs_3[4] 拼接) |
| 动作 action | 离散 1 分支 5 个动作(整数索引 0..4,非连续) |
| 网络 body | hidden_units=512,num_layers=2,激活 SiLU(Swish) |
| 归一化 | normalize=false(无观测归一化统计量) |
| 检查点 | Pyramids/Pyramids-3778906.pt(最终 best brain) |
注:仓库
configuration.yaml配置了vis_encode_type=simple、engine 84x84, 但最终导出图中不含视觉编码器——观测为纯 172 维向量,动作为离散分支。 因此--img仅作为演示入口(把 84x84 图按固定规则合成 172 维向量), 真实推理请使用--obs传入完整 172 维向量观察。
确定性策略重建(与 onnx 完全一致):
x = obs # (172,)
h = silu(x @ W0.T + b0) # (172) -> (512)
h = silu(h @ W1.T + b1) # (512) -> (512)
logits = h @ Wact.T + bact # (512) -> (5)
action = argmax(logits) # 离散动作索引 0..4inference.py # NPU 推理入口(CLI + FastAPI 双模式)
download_weights.py # 权重双源下载(GitCode 镜像 -> hf-mirror,Range 断点续传)
requirements.txt # 运行依赖
assets/ # 验证截图(agent_workflow / npu_device_call / model_result)
model_weights/ # 权重缓存(Pyramids.onnx + Pyramids-3778906.pt + 配置,未上传 Git)torch + torch_npu(版本匹配 CANN)requirements.txt:numpy Pillow torch torch_npu onnx fastapi uvicorn
(onnx 仅用于结构解析/基准核对;onnxruntime 可选,用于 CPU 交叉验证)权重会自动下载并缓存到 ./model_weights/ppo-Pyramids/,优先命中本地缓存,
失败自动降级:GitCode 镜像 https://ai.gitcode.com/hf_mirrors/... → hf-mirror
https://hf-mirror.com/...(含 HTML 预览页识别 + 断点续传)。
# 真实推理:172 维向量观察 -> 离散动作
python3 inference.py --obs "x0,x1,...,x171"
# 缺省演示输入(全零向量)
python3 inference.py
# 演示入口:84x84 图片合成 172 维 obs(仅演示,结果仅供参考)
python3 inference.py --img scene.png
# 指定 NPU 设备
python3 inference.py --obs "..." --device 0输出示例:
[NPU] 已初始化设备 npu:0 (Ascend910_9362)
[权重] 本地缓存命中: .../Pyramids-3778906.pt (8656252 bytes)
[加载] 策略重建完成: obs=172 hidden=512 layers=2 act=discrete[5], 参数量 353797, 设备 npu:0
[结果] 离散动作 action=1 (Backward)
[结果] logits = [-3.324616, 3.924245, -2.875275, 1.638147, -3.989551]
[结果] NaN 检查: False
[JSON] {"model": "Nablaaa/ppo-Pyramids", "device": "npu:0", "action": 1, ...}python3 inference.py --server --port 8084GET /health:
curl http://127.0.0.1:8084/health
# {"status":"ok","device":"npu:0","model":"Nablaaa/ppo-Pyramids"}POST /predict({"obs": [172 维数组]}):
curl -X POST http://127.0.0.1:8084/predict \
-H "Content-Type: application/json" \
-d '{"obs":[0.0,0.1,...,0.0]}'
# {"model":"Nablaaa/ppo-Pyramids","device":"npu:0","action":3,
# "action_name":"TurnRight","action_space":"discrete(5)",
# "logits":[...],"probs":[...],"infer_ms":143.5}.to("npu:<dev>")。try-except:设备不可用 / 初始化失败时给出明确报错。0..4,映射:["Forward","Backward","TurnLeft","TurnRight","Switch"]
(具体语义以 Unity Pyramids 环境的 branch 顺序为准)。