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

Python实现MCTS总是选最后可用走法而非最优解求助

蒙特卡洛树搜索(MCTS)选择最优走法失败问题排查

我用Python 3实现了蒙特卡洛树搜索(MCTS),代码无语法错误,但搜索器始终选择最后一个可用走法,而非最优走法。排查许久未找到原因,推测问题出在最优节点的扩展选择逻辑上。这是我的第一个大型项目,代码可能杂乱,敬请谅解。

示例场景:该游戏需三子连线获胜,当前玩家-1只需在两个-1后落子即可获胜,但程序却选择数组最后一个空位,增加数组长度后依旧如此。

import math
import random
import copy

class Node:  # handles the nodes for the tree
    def __init__(self, game, parent=None):
        self.game = game
        self.parent = parent
        self.children = []
        self.numVisited = 0  # number of times node has been traversed
        self.numWins = 0  # number of accumulated wins
        self.untriedActions = game.getActions()  # list of things to still be searched

    def isExpanded(self):
        return len(self.untriedActions) == 0  # no more actions to try means that it's expanded

    def isTerminal(self):
        return self.game.winner() is not None  # if there is a winner the game is over

    def expand(self):
        assert not self.isExpanded()
        chosenAction = self.untriedActions.pop()  # pick an unused action
        newGame = copy.deepcopy(self.game)  # Use deep copy to ensure new instance
        newGame.applyAction(chosenAction)
        newNode = Node(newGame, parent=self)
        self.children.append(newNode)  # push the new node into the children of this one
        return newNode

    def backpropagate(self, result):  # send things up the tree
        self.numVisited += 1
        # If the current node's player is the same as the result, it's a win for this node
        if result == self.game.player:
            self.numWins += 1  # results are formatted 0 for a loss and 1 for a win
        else:
            self.numWins -= 1

        if self.parent is not None:  # if this is not the root node, then we make sure to go further up
            self.parent.backpropagate(-result)  # Negate the result for the opponent

    def scoreUCT(self):  # calculate the UCB formula for a given node
        c = math.sqrt(2)
        assert self.parent is not None  # we should never ask for the score of the root node
        return (self.numWins / self.numVisited) + c * math.sqrt(math.log(self.parent.numVisited) / self.numVisited)

    def rolloutResult(self):  # plays a random game from the node's state and returns result
        return self.game.randomGame()

    def selectBestChild(self):  # picks the most promising child node by UCB scores
        bestScore = -float('inf')  # Initialize to negative infinity to ensure any score is higher
        bestChild = None

        for child in self.children:  # loop through the children and pick the best one
            score = child.scoreUCT()
            if score > bestScore:
                bestScore = score
                bestChild = child
        return bestChild

    def prettyPrint(self):
        print('=====')
        print(f'State: {self.game.state}')
        print(f'Player: {self.game.player}')
        print(f'Number of visits: {self.numVisited}')
        print(f'Number of wins: {self.numWins}')
        print(f'Children: {len(self.children)}')

        if len(self.children) > 0:
            for child in self.children:
                child.prettyPrint()


class Search:
    def __init__(self, game):
        self.game = game
        self.root = Node(game)

    def iteration(self):
        selectedNode = self.root
        while not selectedNode.isTerminal():
            if not selectedNode.isExpanded():
                selectedNode = selectedNode.expand()
            else:
                selectedNode = selectedNode.selectBestChild()

        result = selectedNode.rolloutResult()
        selectedNode.backpropagate(result)

    def treeSearch(self, numIts):
        for i in range(numIts):
            self.iteration()

        bestScore = -float('inf')
        bestChild = None
        for child in self.root.children:  # loop through the children and pick the best one
            score = child.numVisited
            if score > bestScore:
                bestScore = score
                bestChild = child
        return bestChild

    def treePrint(self):
        self.root.prettyPrint()


class Game:  # game logic
    def __init__(self, state, player=1):
        self.state = state[:]
        self.player = player

    def getActions(self):  # return a list of all actions possible from current game state
        if self.winner() is not None:
            return []

        possibleActions = []
        for i in range(len(self.state)):
            if self.state[i] == 0:
                possibleActions.append(i)
        return possibleActions

    def applyAction(self, action):  # apply a given action to the game to change the state.
        assert self.state[action] == 0
        self.state[action] = self.player
        # V THIS IS IMPORTANT
        self.player *= -1

    def randomGame(self):  # play a random game from the current position
        copyGame = copy.deepcopy(self)  # Use deep copy to ensure new instance
        while copyGame.winner() is None:  # play until the game finishes
            possMoves = copyGame.getActions()
            copyGame.applyAction(random.choice(possMoves))
        return copyGame.winner()

    def winner(self):
        for i in range(0, len(self.state) - 2):
            if self.state[i] != 0 and self.state[i] == self.state[i + 1] == self.state[i + 2]:
                return self.state[i]
        if 0 not in self.state:
            return 1

        return None


# Example game and search
# This game needs a 3-in-a-row to win, and so playing after the two -1s would be a win for -1. 
# However, this will play the final "square" in the array
# If you add more rows to the array, it will still choose the last one.
g = Game([0, 0, 0, 1, -1, -1, 0, 0], -1)
s = Search(g)
best_child = s.treeSearch(100)
print(best_child.game.state)

问题修复方案

1. 修正平局判定逻辑

在Game.winner()中,平局时返回0而非1,避免错误判定玩家1获胜:

def winner(self):
    for i in range(0, len(self.state) - 2):
        if self.state[i] != 0 and self.state[i] == self.state[i + 1] == self.state[i + 2]:
            return self.state[i]
    if 0 not in self.state:
        return 0  # 平局返回0

    return None

2. 修复回溯计分逻辑

回溯时,根据当前节点玩家与结果的关系正确计分,父节点无需翻转结果,而是基于自身玩家判断:

def backpropagate(self, result):
    self.numVisited += 1
    # 结果为0是平局,不增减分数;结果为玩家编号则对应胜负
    if result == self.game.player:
        self.numWins += 1
    elif result != 0:
        self.numWins -= 1

    if self.parent is not None:
        self.parent.backpropagate(result)  # 父节点用原始结果判断自身分数

3. 扩展时随机选择未尝试动作

避免总是从列表末尾取动作,确保初始扩展的公平性:

def expand(self):
    assert not self.isExpanded()
    # 随机选择一个未尝试的动作
    chosenAction = random.choice(self.untriedActions)
    self.untriedActions.remove(chosenAction)
    newGame = copy.deepcopy(self.game)
    newGame.applyAction(chosenAction)
    newNode = Node(newGame, parent=self)
    self.children.append(newNode)
    return newNode

4. 根节点选择最优子节点的正确逻辑

根节点应选择胜率最高的子节点,而非访问次数最多的:

def treeSearch(self, numIts):
    for i in range(numIts):
        self.iteration()

    bestScore = -float('inf')
    bestChild = None
    for child in self.root.children:
        # 计算胜率,避免除以0(初始访问次数为0的情况)
        win_rate = child.numWins / child.numVisited if child.numVisited > 0 else -float('inf')
        if win_rate > bestScore:
            bestScore = win_rate
            bestChild = child
    return bestChild

内容的提问来源于stack exchange,提问作者Elliott

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 01:10:58