You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 06:35:25