Unity斯诺克游戏中ML-Agents实现失败,请求问题排查
问题描述
尝试在Unity斯诺克游戏中实现基于ML-Agents的强化学习玩家,但训练时Agent完全没有动作,需排查代码错误并给出解决方案。
核心代码
using System; using Unity.MLAgents; using Unity.MLAgents.Sensors; using Unity.MLAgents.Actuators; using UnityEngine; using System.Collections; using UnityEngine.UI; public struct Player { public int score; public bool isNPC; } public class Snooker2D : Agent { public GameObject cueBall; public Transform stick; public Transform ball; public Transform selectionFx; public Slider slider; int currentPlayer = 1; int nPlayers = 2; Player[] players = new Player[2]; bool follow = true; bool doubleShot = false; float ballRadius = 0; LayerMask layerMaskBalls = 1 << 9; LayerMask layerMaskBallsAndWalls = 1 << 9 | 1 << 10; LayerMask layerMaskWalls = 1 << 10; RaycastHit2D hit; float dist = 0; float minDist = 0; float maxDist = -3; float forceMultiplier = 2.5f; Vector3 whiteBallPos = new Vector3(-4.3111f, 0, 0); public override void OnEpisodeBegin() { // Reset the environment Application.LoadLevel(Application.loadedLevel); } public override void CollectObservations(VectorSensor sensor) { // Target and Agent positions sensor.AddObservation(ball.position); // sensor.AddObservation(stick.position); // All balls positions foreach (GameObject ball in GameObject.FindGameObjectsWithTag("Ball")) { sensor.AddObservation(ball.transform.position); } } public override void OnActionReceived(ActionBuffers actionBuffers) { // Actions, size = 3, the force vector var action = actionBuffers.ContinuousActions; Vector3 force = new Vector3(action[0], action[1], action[2]); // Apply the force vector to the game environment, for example by adding it to the velocity of the cue ball Shoot(force); } public override void Heuristic(in ActionBuffers actionsOut) { var continuousActionsOut = actionsOut.ContinuousActions; // Set the force vector using keyboard or gamepad inputs, each action is between -1 and 1 // Each action is an axis of the force vector continuousActionsOut[0] = Input.GetAxis("Horizontal"); continuousActionsOut[1] = Input.GetAxis("Vertical"); continuousActionsOut[2] = Input.GetAxis("Jump"); } void Start() { players[1].isNPC = true; ballRadius = ball.GetComponent<CircleCollider2D>().radius; minDist = -(ballRadius + ballRadius / 2); dist = Mathf.Clamp(maxDist / 2, maxDist, minDist); slider.maxValue = -maxDist; slider.minValue = -minDist; slider.value = dist + -maxDist - minDist; selectionFx.GetComponent<Fader>().StartFade(); } void Update() { if (Input.GetKeyDown("r")) Application.LoadLevel(Application.loadedLevel); // PLAYER TURN if (follow) { Vector3 forceDir = Vector3.zero; // Analog Player if (!players[currentPlayer].isNPC) { // MOUSE: Get mouse position Vector3 mPos = Camera.main.ScreenToWorldPoint(new Vector3(Input.mousePosition.x, Input.mousePosition.y, 10)); PowerAdjust(); RotateStickAroundBall(mPos); ProjectTrajectory(); forceDir = (ball.position - stick.position).normalized * -(dist - minDist - 0.02f) * forceMultiplier; if (Input.GetMouseButtonUp(0)) Shoot(forceDir); // && hit.collider != null) } // AI Player else { // TODO: Here will be the AI // Reward the model if there are less balls on the table // Apply the action to the ball //forceDir = new Vector3(1, 1, 1); //Shoot(forceDir); var continuousActionsOut = new ActionBuffers(); Heuristic(continuousActionsOut); SetReward(1.0f / GameObject.FindGameObjectsWithTag("Ball").Length); if (GameObject.FindGameObjectsWithTag("Ball").Length == 1) { // End the episode if there is only one ball left EndEpisode(); } } } // start following if (!follow) { if (AllBallsStopped() && ball.GetComponent<Rigidbody2D>().velocity.sqrMagnitude == 0.0f) { // Update player score // players[currentPlayer].score = UpdateScore(currentPlayer); // currentPlayer = (currentPlayer + 1) % nPlayers; follow = true; if (cueBall.activeSelf == false) { SetReward(-0.1f); RespawnBall(); } selectionFx.GetComponent<Fader>().StartFade(); Invoke("HideShowStick", 0.2f); } } } // update() }
代码错误分析
- 未触发ML-Agents决策流程:AI分支中没有调用
RequestDecision(),ML-Agents需要Agent主动请求决策才会调用OnActionReceived执行动作。 - 场景重置方式错误:使用
Application.LoadLevel重载场景是旧API,会破坏ML-Agents的训练上下文,导致训练状态丢失。 - 观测空间不规范:
- 每次调用
CollectObservations都通过GameObject.FindGameObjectsWithTag查找球,性能差且观测数量随球被打进而变化,不符合ML-Agents对固定维度观测的要求。 - 未对观测值进行归一化,ML-Agents模型对未归一化的输入收敛速度极慢甚至无法收敛。
- 每次调用
- 动作空间冗余:2D斯诺克游戏中Z轴力无效,当前3维动作空间包含多余维度,增加了模型学习难度。
- 奖励机制不合理:
- 奖励在AI分支每帧发放,而非根据动作结果(如打进球、母球犯规)在回合结束后计算,无法正确引导模型学习。
- 奖励值设置缺乏区分度,无法有效激励正确行为、惩罚错误行为。
- 错误调用Heuristic方法:
Heuristic是用于人类手动控制的方法,AI分支手动调用它无法让ML-Agents的决策系统生效。 - 状态控制逻辑缺陷:AI射击后未设置
follow = false,导致一直停留在AI分支循环,无法进入球运动后的状态判断流程。
解决方案
1. 触发ML-Agents决策流程
修改AI分支代码,添加决策请求并更新回合状态:
// AI Player else { if (follow) { RequestDecision(); follow = false; // 射击后进入等待球停止状态 } }
2. 修复场景重置逻辑
替换场景重载,手动重置所有球的位置、速度和状态,需提前缓存球的初始位置:
private List<GameObject> allBalls; private Dictionary<GameObject, Vector3> ballInitialPositions; void Start() { players[1].isNPC = true; ballRadius = ball.GetComponent<CircleCollider2D>().radius; minDist = -(ballRadius + ballRadius / 2); dist = Mathf.Clamp(maxDist / 2, maxDist, minDist); slider.maxValue = -maxDist; slider.minValue = -minDist; slider.value = dist + -maxDist - minDist; selectionFx.GetComponent<Fader>().StartFade(); // 缓存所有球及其初始位置 allBalls = new List<GameObject>(GameObject.FindGameObjectsWithTag("Ball")); ballInitialPositions = new Dictionary<GameObject, Vector3>(); foreach (var b in allBalls) { ballInitialPositions.Add(b, b.transform.position); } } public override void OnEpisodeBegin() { // 重置母球 cueBall.transform.position = whiteBallPos; Rigidbody2D cueRb = cueBall.GetComponent<Rigidbody2D>(); cueRb.velocity = Vector2.zero; cueRb.angularVelocity = 0; cueBall.SetActive(true); // 重置其他球 foreach (var b in allBalls) { b.transform.position = ballInitialPositions[b]; Rigidbody2D rb = b.GetComponent<Rigidbody2D>(); rb.velocity = Vector2.zero; rb.angularVelocity = 0; b.SetActive(true); } // 重置回合状态 follow = true; currentPlayer = 1; }
3. 优化观测空间
使用缓存的球引用,将观测值归一化为相对母球的坐标:
public override void CollectObservations(VectorSensor sensor) { Vector3 cuePos = cueBall.transform.position; Rigidbody2D cueRb = cueBall.GetComponent<Rigidbody2D>(); // 添加母球的速度(归一化) sensor.AddObservation(cueRb.velocity.normalized); // 添加所有球的相对位置(归一化到[-1,1]) foreach (var b in allBalls) { if (b.activeSelf) { Vector3 relativePos = b.transform.position - cuePos; // 假设场景最大范围为10,缩放观测值 sensor.AddObservation(relativePos / 10f); } else { // 球被打进后用固定值标记 sensor.AddObservation(new Vector3(2f, 2f, 0f)); } } }
4. 调整动作空间
在Unity编辑器的Agent组件中,将Continuous Actions数量改为3(X方向、Y方向、力度),然后修改OnActionReceived:
public override void OnActionReceived(ActionBuffers actionBuffers) { var actions = actionBuffers.ContinuousActions; // 提取方向并归一化 Vector2 direction = new Vector2(actions[0], actions[1]).normalized; // 提取力度并限制在0-1范围 float force = Mathf.Clamp01(actions[2]) * forceMultiplier; // 给母球施加冲量 Rigidbody2D cueRb = cueBall.GetComponent<Rigidbody2D>(); cueRb.AddForce(direction * force, ForceMode2D.Impulse); }
5. 完善奖励机制
在球停止运动后根据结果计算奖励:
void Update() { // ... 原有代码 ... // 球运动后的状态处理 if (!follow) { if (AllBallsStopped() && cueBall.GetComponent<Rigidbody2D>().velocity.sqrMagnitude < 0.01f) { int remainingBalls = GameObject.FindGameObjectsWithTag("Ball").Length; int ballsHitIn = allBalls.Count - remainingBalls; // 打进球给予正奖励 if (ballsHitIn > 0) { SetReward(ballsHitIn * 2f); } // 未打进球给予小负奖励 else { SetReward(-0.1f); } // 母球进洞给予较大负奖励 if (!cueBall.activeSelf) { SetReward(-1f); RespawnBall(); } // 检查是否结束回合 follow = true; if (remainingBalls == 1) { EndEpisode(); } } } }
6. 移除错误的Heuristic调用
删除AI分支中手动调用Heuristic的代码,让ML-Agents自动处理决策流程。
7. 补充缺失方法
确保Shoot、AllBallsStopped、RespawnBall等方法正确实现:
void Shoot(Vector2 force) { Rigidbody2D cueRb = cueBall.GetComponent<Rigidbody2D>(); cueRb.AddForce(force, ForceMode2D.Impulse); } bool AllBallsStopped() { foreach (var b in allBalls) { if (b.activeSelf && b.GetComponent<Rigidbody2D>().velocity.sqrMagnitude > 0.01f) { return false; } } return true; } void RespawnBall() { cueBall.transform.position = whiteBallPos; cueBall.GetComponent<Rigidbody2D>().velocity = Vector2.zero; cueBall.SetActive(true); }
内容的提问来源于stack exchange,提问作者Pablo Olivares
相关产品推荐
相关产品推荐

