w
gcw_uQ09W7jl/HumanCompatibleAI-ppo-seals-CartPole-v0-NPU
模型介绍
文件和版本
Pull Requests
讨论
分析

HumanCompatibleAI/ppo-seals-CartPole-v0 NPU 适配

模型任务与契约

  • 任务:离散强化学习 seals/CartPole-v0 倒立摆平衡控制,观测为 4 维连续向量 Box(4),输出为 2 离散动作的贪心决策。
  • 架构:PPO MlpPolicy 纯 PyTorch 4 ->64 ReLU ->64 ReLU ->2 logits,9155 参数(含 value 头,policy-only 4610),float32,Discrete(2),无归一化,ReLU 激活,net_arch pi [64,64] vf [64,64],SB3 兼容权重直接加载。
  • 观测 schema:obs [B,4] float32,维度 [cart_pos, cart_vel, pole_angle, pole_ang_vel],本次固定 batch=8 种子 0,hash efc40145060ba0d4,示例 obs0 [0,0,0,0] obs1 [0.5,0.3,0.1,-0.2] obs2 [-0.5,-0.3,-0.1,0.2]。
  • 动作 schema:logits [B,2] float32 → softmax -> probs [B,2] → argmax -> deterministic [B] in {0,1},horizon=1 decision,action_dim=2,对应 seals/CartPole-v0 离散 2 动作(0 左移, 1 右移);本次 det [0,1,0,1,0,1,0,1]。
  • 权重:https://huggingface.co/HumanCompatibleAI/ppo-seals-CartPole-v0 ,revision a86541dd3744275227cb16d07192967fc6b7775e,policy.pth 40641 bytes + data,约 9.2K 参数,本仓通过本地 /tmp/ppo-seals-CartPole-v0/ppo-seals-CartPole-v0/policy.pth 加载,兼容 snapshot_download HF_ENDPOINT https://hf-mirror.com,训练超参 gamma 0.9999 batch 256 n_steps 512 n_epochs 10 lr 0.00124 ent_coef 0.0085 训练 100k steps。
  • 环境:seals/CartPole-v0 (gymnasium + seals wrapper),与 CartPole-v0 一致但带 seals 插件,reward 均值 500 +/-0 deterministic 10 episodes。

真实结果

  • 设备:Ascend910_9362,npu:0,torch_npu 2.9.0.post1+gitee7ba04,torch 2.9.0+cpu,CANN 8.5.1,2 cards 使用 npu:0,npu-smi 正常,torch.npu.is_available() true。
  • 输入:obs [8,4] float32 hash efc40145060ba0d4,obs_npu npu:0 float32,obs_cpu cpu float32,首参 npu:0 float32 [64,4] 与 cpu float32 [64,4] 对照,hidden/logits/actions 均在 npu:0。
  • 输出:
    • CPU logits [8,2] batch0 [0.511556, -0.509789] probs [0.73523, 0.264766] batch1 [-1.84245, 1.84460] probs [0.02443, 0.97557] batch3 [-4.23542, 4.23841] det [0,1,0,1,0,1,0,1] device cpu。
    • NPU 同形状 [8,2] npu:0 float32 batch0 [0.511556, -0.509789] batch1 [-1.84245, 1.84460] batch3 [-4.23543, 4.23841] det [0,1,0,1,0,1,0,1] device npu:0,hidden npu:0 logits npu:0 actions npu:0,无 CPU fallback。
    • 有限值 PASS,shape 一致 [8,2],动作在 [0,1] bounds 内。
  • 一致性:CPU vs NPU 同权重同观测 logits max_abs 4.768e-07 mean_abs 1.90e-07 probs max_abs 7.45e-08,tolerance 0.004338 (atol 1e-4 + rtol 1e-3 * max_ref 4.238),det 完全相等,PASS,validate_policy_outputs.py count 16 finite true passed true。
  • 性能:首轮编译 133.07 ms(含图编译),warmup 3,稳定 avg 0.928 ms min 0.906 max 0.938 p50 0.929 p90 0.934 p95 0.938(10 runs),~8618 decisions/s / ~8618 actions/s,batch 8, obs 4, horizon 1, float32,torch.npu.synchronize() 包围计时,吞吐 decisions_per_s = batch*1000/avg。

CPU-NPU 一致性

  • 共享权重(同 policy.pth 4610 policy 参数)、观测(同 batch 8 固定向量 hash efc40145060ba0d4)、dtype(float32)、采样(deterministic argmax)、seed(0 固定观测)、evaluation mode。
  • validate_policy_outputs.py --atol 1e-4 --rtol 1e-3 阈值 FP32 默认,真实 max_abs 4.76e-07 远小于阈值 0.004338。
  • 设备证明:first_param_device npu:0(CPU 对照为 cpu),obs_npu npu:0,hidden npu:0,logits npu:0,actions npu:0,无 .cpu() 隐式路径,纯 PyTorch 实现无 SB3 隐式 CPU。

性能

  • 首轮 133 ms,稳态 0.928 ms,吞吐 8618 decisions/s。4.6K MLP 在 NPU 上极低延迟,适合实时控制。benchmark_policy.py avg_ms 0.928 decisions_per_s 8618。

环境限制

  • 单卡 npu:0 即可,模型 4.6K 参数 <1MB,显存 <100MB,远低于 64GB HBM。
  • 依赖:torch 2.9, torch_npu 2.9.post1, numpy, huggingface_hub,无需 stable-baselines3/gymnasium/seals 即可推理(权重兼容 SB3)。
  • 已知:seals/CartPole-v0 环境 Box(4) Discrete(2),训练 100k steps,batch 256,gamma 0.9999;推理仅验证离散 logits/probs/argmax 而非闭环 Gym rollout;环境 reward 500 +/-0 为额外评估,不替代数值一致性。

复现

python inference.py
# 自动解析 /tmp/ppo-seals-CartPole-v0/ppo-seals-CartPole-v0/policy.pth 或 HF cache 或 snapshot_download
# 打印模型/revision/route/backend/device/dtype、输入摘要、logits/det 摘要、同步耗时、PASS/FAIL

证据图

  • agent workflow
  • npu device call
  • model result

仓库结构

inference.py
readme.md
requirements.txt
assets/agent_workflow.png
assets/npu_device_call.png
assets/model_result.png

许可证 MIT;权重遵循原 HumanCompatibleAI/ppo-seals-CartPole-v0。