尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

用强化学习训练「Plane Strike」棋类智能体:TF Agents / TensorFlow / JAX 三路径实战与 TFLite 端侧部署指南

用强化学习训练「Plane Strike」棋类智能体:TF Agents / TensorFlow / JAX 三路径实战与 TFLite 端侧部署指南 示例工程【免费下载链接】examplesTensorFlow examples项目地址https://gitcode.com/gh_mirrors/exam/examples点击查看免费下载本文基于 TensorFlow examples 仓库中的 reinforcement_learning 示例训练代码位于 ml 目录完整讲解如何从零训练一个能在手机上对弈的强化学习智能体。读者将掌握Plane Strike 棋类环境的建模方式、TF Agents / TensorFlow / JAX 三条独立训练路径的完整操作步骤、训练监控方法以及如何把训练好的策略转换为 TFLite 模型并集成到 Android 应用中。背景Plane Strike 是什么Plane Strike 是一个回合制棋类游戏玩法与经典的 Battleship海战棋类似二者的核心差异在于Battleship 允许你在棋盘上摆放多艘长度 2–5 格的战舰而 Plane Strike 每方只在开局时摆放一架飞机。游戏界面包含上下两块 8×8 棋盘上方是智能体Agent的棋盘下方是玩家自己的棋盘。开局时双方飞机位置随机生成并互相隐藏你需要先于对方找出其飞机的全部 8 个格子在己方棋盘上以 8 个蓝色格子展示若不满意可重置游戏重新随机摆放先找齐对方全部飞机格的一方获胜随后游戏自动重新开始。对局中玩家点击上方智能体棋盘中的某个格子若该格是飞机格则变红命中 hit否则变黄未命中 miss。应用会同时统计双方棋盘上的命中次数方便你掌握战况。重复点击同一格子只会浪费出手机会不会产生任何效果。这套玩法天然适合用强化学习建模智能体看到的只有自己不断试错得到的部分信息己方棋盘上记录的命中/未命中结果需要在有限步数内推理出对方飞机的完整布局是一个典型的序贯决策问题。训练代码目录结构训练相关代码集中在 lite/examples/reinforcement_learning/ml 目录下common.pyTF/JAX 两条路径共用的工具函数包括棋盘尺寸常量、折扣因子、随机飞机摆放逻辑、对弈数据收集与折扣奖励计算tf_agents/TF Agents 路径包含自定义 PyEnvironment 环境、训练脚本与对应依赖清单tf_and_jax/TensorFlow 与 JAX 两条路径内含一个标准 OpenAI Gym 环境通过setup.py安装Android 端工程位于 lite/examples/reinforcement_learning/android其 assets 目录 已预置了两种已训练好的 TFLite 模型planestrike_tf.tflite来自 TF/JAX 路径与planestrike_tf_agents.tflite来自 TF Agents 路径。三条训练路径概览仓库提供了三种可选的训练实现你可以任选其一自行训练模型路径实现框架环境形式备注TF Agentstf_agents/training_tf_agents.py自定义PyEnvironmentplanestrike_py_environment.py使用 REINFORCE Agent Reverb 回放缓冲区输出 4 输入 TFLite 模型TensorFlowtf_and_jax/training_tf.pyOpenAI Gym 环境gym_planestrike手写 REINFORCE 风格训练循环输出单输入 TFLite 模型JAX / Flaxtf_and_jax/training_jax.pyOpenAI Gym 环境gym_planestrike实验性实现README 明确标注highly experimental同时训练效率与效果还有进一步的提升空间例如利用棋盘的对称性旋转/镜像等价局面、更精细的奖励塑形reward shaping以及并行采样运行等。环境准备与依赖安装开始训练前先安装对应路径的依赖。TF Agents 路径的依赖在 tf_agents/requirements.txt内容为tensorflow2.7.2 tf-agents0.9.0 dm-reverb0.5.0 tensorflow-probability0.14.0 Pillow10.0.1TF/JAX 路径的依赖在 tf_and_jax/requirements.txt内容为flax0.6.0 jax0.3.17 jaxlib0.3.15 tensorboard2.9.1 tensorflow2.9.1 numpy1.23.2 gym0.18.0 tqdm4.64.0安装命令统一为pip install -r requirements.txt需要注意README 撰写时2021 年 7 月只有 nightly 版本的 TF / TF Agents 支持将策略直接转换为 TFLite 模型仓库当前通过requirements.txt固定了上述正式版本号训练与转换流程以仓库内实际锁定的版本组合为准。若使用其他版本建议先验证TFLiteConverter.from_saved_model(..., signature_keys[action])的转换链路是否可用。路径一基于 TF Agents 训练并转换策略TF Agents 路径提供了最工业级的训练体验用 Reverb 做回放缓冲区、用py_driver收集对局、用 REINFORCE Agent 做策略梯度更新并在训练结束后直接把 SavedModel 策略转换为 TFLite。1. 自定义 PyEnvironmentPlaneStrikePyEnvironmentTF Agents 环境定义在 planestrike_py_environment.py它继承py_environment.PyEnvironment核心要素如下观察空间(8, 8)的float32棋盘矩阵取值上限 1命中、下限 -1未命中、0 表示未试探即智能体看到的是一张只含自身打击结果的可见棋盘动作空间BoundedArraySpec约束的int32标量范围0 ~ board_size**2 - 1即把 8×8 棋盘按行优先展平后的格子编号奖励设计源码中明确定义命中HIT_REWARD 1未命中MISS_REWARD 0重复打击已试探格子REPEAT_STRIKE_REWARD -1在max_steps默认BOARD_SIZE**2 64步内找齐全部飞机格FINISHED_GAME_REWARD 10超时未完成UNFINISHED_GAME_REWARD -10终止条件命中数达到飞机大小8或步数耗尽回合自动结束并重置。棋盘的随机初始化逻辑位于 common.py 的initialize_random_hidden_board先随机选择飞机朝向右/上/左/下四种方向再确定机身十字核心位置最后填充尾翼三格合计 8 个占用格该文件同时定义了BOARD_SIZE 8、PLANE_SIZE 8与折扣因子GAMMA 0.5。2. 训练脚本与超参数在tf_agents目录下执行python training_tf_agents.py tensorboard --logdir./tf_agents_logtraining_tf_agents.py 的关键配置ITERATIONS 250000总训练迭代数COLLECT_EPISODES_PER_ITERATION 1每轮用 collect policy 采集 1 局对局存入回放缓冲区REPLAY_BUFFER_CAPACITY 2000Reverb 回放缓冲区容量DISCOUNT 0.5折扣因子FC_LAYER_PARAMS 100actor 网络隐藏层宽度InnerReshape展平后接两层 Dense最后以Categorical分布输出 64 维动作 logitsLEARNING_RATE 1e-3Adam 优化器学习率NUM_EVAL_EPISODES 20、EVAL_INTERVAL 500每 500 次迭代用 20 局评估平均回报与平均对局长度并写入 TensorBoard回放缓冲区使用 Reverb 的uniform_tableuniform 采样 FIFO 淘汰MinSize(1)限速。训练过程中 TensorBoard 会记录两条标量曲线Average return平均回报整体应随训练上升与Average episode length平均对局长度整体应随训练下降。看到智能体变聪明的标志正是平均回报上升、平均对局长度下降。3. 将策略导出为 TFLite 模型训练完成后脚本先把 collect policy 保存为 SavedModel./policy目录再通过以下代码转换为 TFLiteconverter tf.lite.TFLiteConverter.from_saved_model( policydir, signature_keys[action]) converter.target_spec.supported_ops [ tf.lite.OpsSet.TFLITE_BUILTINS, # 启用 TensorFlow Lite 内置算子 tf.lite.OpsSet.SELECT_TF_OPS # 启用完整 TensorFlow 算子flex 模式 ] tflite_policy converter.convert()最终产物为planestrike_tf_agents.tflite。与 TF/JAX 路径产出的模型不同TF Agents 转换出的 TFLite 模型带有 4 个输入张量策略签名中携带了step_type、discount、observation、reward等时间步信息其中真正对推理有用的只有第 3 个输入observation形状[1, 8, 8]的 float32 棋盘观测输出为[1]的 int32 动作索引其结构如下图所示路径二TensorFlow 直接训练REINFORCE 风格Step 1. 安装 OpenAI Gym 环境进入tf_and_jax/gym_planestrike目录执行python setup.py install该目录是一个标准的 Python 包工程含setup.py安装后环境注册名为PlaneStrike-v0可在 gym_planestrike/envs/planestrike.py 中查看其reset/step实现细节其奖励与终止语义与 TF Agents 路径的环境保持一致。Step 2. 训练模型在tf_and_jax目录下执行python training_tf.py tensorboard --logdir./tf_logtraining_tf.py 是一个简洁的手写策略梯度实现网络为三层全连接Flatten(8×8)→Dense(128, relu)→Dense(64, relu)→Dense(64, softmax)即以观测棋盘为输入、输出 64 个格子的打击概率分布优化器使用 SGDLEARNING_RATE 0.002损失为sparse_categorical_crossentropy用每步的折扣奖励作为sample_weight加权——这正是 REINFORCE 策略梯度带折扣回报的经典写法ITERATIONS 80000每轮调用 common.py 中的play_game打一局并记录棋盘、动作、奖励序列再由compute_rewards从终局向前回溯计算折扣累计奖励默认gamma 0.5训练期间记录game_length每局步数标量README 建议在 TensorBoard 中对该曲线使用0.9 的平滑系数smoothing factor观察趋势。TensorFlow 路径的game_length训练曲线大致呈现随步数增加持续下降并收敛的形态训练前期下降明显后期趋于平稳说明智能体用更少的步数找齐飞机格训练曲线示例如下Step 3. 导出 TFLite 模型训练结束后脚本调用TFLiteConverter.from_keras_model(model)把 Keras 模型直接转换为planestrike.tflite位于当前目录。该模型只有 1 个输入[1, 8, 8]的观测棋盘推理接口比 TF Agents 产物更简单。路径三JAX / Flax实验性同样在tf_and_jax目录下执行python training_jax.py tensorboard --logdir./jax_logtraining_jax.py 使用 JAX / Flax 实现了同一套 REINFORCE 训练流程监控指标同样是game_length其收敛行为与 TF 路径一致步数增加、每局长度下降训练曲线示例见下图。需要留意 README 的明确提示JAX 实现属于高度实验性highly experimental代码适合对 JAX/Flax 生态感兴趣的读者参考生产环境建议优先选择 TF Agents 或 TensorFlow 路径。把模型集成进 Android 应用训练得到的模型文件需要放进 Android 工程的 assets 目录lite/examples/reinforcement_learning/android/app/src/main/assetsTF Agents 路径产出的planestrike_tf_agents.tflite可直接使用与预置文件名一致TF/JAX 路径产出的planestrike.tflite需要重命名为planestrike_tf.tflite替换 assets 目录中现有的同名文件。Android 端对应的推理封装位于 org/tensorflow/lite/examples/reinforcementlearning 包下其中PlaneStrikeAgent.java、RLAgent.java负责加载并运行单输入的 TF 模型RLAgentFromTFAgents.java负责处理 TF Agents 产物的 4 输入模型在调用时需额外提供step_type、discount、reward等占位张量仅observation携带真实棋盘数据。构建运行步骤可参考 android/README.md用 Android Studio 4.2 导入reinforcement_learning/android工程连接真机或模拟器后Run - Run app即可在设备上以Reinforcement Learning应用与训练出的智能体对弈。提升训练效果的方向README 明确指出训练速度和效果仍有较大优化空间主要方向包括利用棋盘对称性Plane Strike 棋盘存在旋转/镜像对称的等价局面可对观测做数据增强或让网络共享对称权重减少冗余探索奖励塑形reward shaping当前奖励设计是稀疏的命中 1、未命中 0、重复 -1、终局 ±10可考虑引入距离启发、命中概率估计等稠密奖励加速收敛并行采样TF Agents 路径可增大COLLECT_EPISODES_PER_ITERATION或使用分布式收集JAX 路径则天然适合vmap/pmap并行打多局对局。这些改造都发生在训练侧模型导出与 Android 推理链路无需变动。参考资料OpenAI Spinning Up 中的 VPGVanilla Policy Gradient算法文档本仓库 TF/JAX 路径采用的正是这类策略梯度思想Deep reinforcement learning 与 battleship 的工程实践文章讨论了用强化学习求解海战棋类不完全信息博弈的思路与 Plane Strike 的建模方式同源。文中涉及的全部训练脚本、环境实现与依赖清单均可直接查阅仓库源码TF Agents 环境与训练见 tf_agents 目录TF/JAX 训练见 tf_and_jax 目录公共工具函数见 common.pyAndroid 推理与集成见 android 目录。赞分享示例工程【免费下载链接】examplesTensorFlow examples项目地址https://gitcode.com/gh_mirrors/exam/examples点击查看免费下载相关推荐AReaL 智能体强化学习Agentic RL实战指南用 OpenAI Agents、CAMEL-AI 与 LangChain 训练可部署的智能体AReaL 智能体强化学习Agentic RL实战指南用 OpenAI Agents、CAMEL AI 与 LangChain 训练可部署的智能体 ARe人工智能大模型强化学习分布式训练AI AgentDopamine JAX 版使用指南JAX 加速的强化学习智能体架构与训练实践Dopamine JAX 版使用指南JAX 加速的强化学习智能体架构与训练实践 Dopamine 是一个面向强化学习算法快速原型验证的研究框架而 dopam强化学习机器学习深度学习Google AI Edge Gallery在手机上完整跑起大语言模型的实操指南Google AI Edge Gallery在手机上完整跑起大语言模型的实操指南 Google AI Edge Gallery 是 Google AI Edg示例工程上一篇AppleALC为Hackintosh解锁原生macOS音频体验的完整指南下一篇Halo 个人中心附件上传存储策略问题分析与解决方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表