模型名称: decision-transformer-gym-hopper-medium(Decision Transformer) 模型链接: edbeeching/decision-transformer-gym-hopper-medium 模型描述: Decision Transformer(Chen et al., 2021)将强化学习建模为序列生成问题:把 Return-to-Go(期望回报)、状态、动作三种模态交替拼成序列,输入 GPT-2 风格的因果 Transformer,自回归地预测下一时刻的最优动作。推理时通过调节目标回报即可控制策略行为,无需价值函数与策略梯度训练。本 checkpoint 在 MuJoCo Hopper-v3 的 D4RL medium(半专家)数据集上离线训练。 模型架构: GPT-2 因果 Transformer 骨干 + 状态/回报/动作线性嵌入 + 时间步 embedding,3 层 decoder(n_embd=768, n_head=1, hidden_size=128),state_dim=11, act_dim=3,动作输出经 tanh 限幅。 参数规模: 约 0.86M(pytorch_model.bin 约 6.3 MB,fp32) 论文: Decision Transformer: Reinforcement Learning via Sequence Modeling (NeurIPS 2021)
| 依赖项 | 版本要求 | 说明 |
|---|---|---|
| Python | >= 3.10 | 推荐 3.11 |
| torch / torch_npu | 2.1.0+ | torch_npu 从 CANN 安装,使用共享环境 |
| transformers | >= 4.40 | 内置 DecisionTransformerModel |
| huggingface_hub | < 1.0 | 与 transformers 版本兼容 |
| numpy | >= 1.24 | 数值计算 |
安装命令:
pip install --target /mnt/workspace/decision-transformer-gym-hopper-medium/libs -i https://pypi.tuna.tsinghua.edu.cn/simple \
"transformers>=4.40" "huggingface_hub<1.0" numpy
# torch/torch_npu 使用共享环境# 检查 NPU 设备
npu-smi info
# 权重位于 /home/developer/models/decision-transformer-gym-hopper-medium(config.json + pytorch_model.bin)PYTHONPATH=/mnt/workspace/decision-transformer-gym-hopper-medium/libs \
python3 inference.py --model_path /home/developer/models/decision-transformer-gym-hopper-medium脚本自动选择设备 npu > cuda > cpu,内置 Hopper 示例轨迹(长度 30),无需其他必填参数。
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
| --model_path | str | /home/developer/models/decision-transformer-gym-hopper-medium | 本地权重目录 |
| --device | str | 自动(npu>cuda>cpu) | 推理设备 |
| --target_return | float | 600.0 | 目标回报(Return-to-Go 初始值),控制策略激进程度 |
输入:
长度 30 的合成 Hopper 轨迹:11 维状态随机游走,returns-to-go 从 600 线性衰减,timesteps 递增输出:
[INFO] device = npu
[INFO] 模型加载完成 load_time=x.xx s params=0.86M (...)
[INFO] 动作预测 shape=(3,), dtype=float32, values=[...]
[RESULT] status=ok device=npu 耗时/elapsed=x.xxx s


DecisionTransformerModel 架构,通过 from_pretrained 加载本地权重目录并设置 local_files_only=True,全程离线,不联网下载;/mnt/workspace/decision-transformer-gym-hopper-medium/libs,运行时由脚本自动叠加 PYTHONPATH,不污染共享环境;共享环境的 huggingface_hub>=1.0 与 transformers 不兼容,libs 内已固定 huggingface_hub<1.0 并优先加载;DecisionTransformerConfig.from_pretrained 会自动忽略,不影响加载。