From 0ac37a33353fc5dcd81a13a5e23ad1bcf93dd0d3 Mon Sep 17 00:00:00 2001 From: Lambda Date: Tue, 1 Dec 2020 01:30:17 +0800 Subject: [PATCH 01/12] clearify README --- README.md | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) mode change 100644 => 100755 README.md diff --git a/README.md b/README.md old mode 100644 new mode 100755 index 0d160e7..dc2ae78 --- a/README.md +++ b/README.md @@ -1,8 +1,8 @@ # MCTS -This package provides a simple way of using Monte Carlo Tree Search in any perfect information domain. +This package provides a simple way of using Monte Carlo Tree Search in any perfect information domain. -## Installation +## Installation With pip: `pip install mcts` @@ -10,15 +10,15 @@ Without pip: Download the zip/tar.gz file of the [latest release](https://github ## Quick Usage -In order to run MCTS, you must implement a `State` class which can fully describe the state of the world. It must also implement four methods: +In order to run MCTS, you must implement a `State` class which can fully describe the state of the world. It must also implement four methods: - `getCurrentPlayer()`: Returns 1 if it is the maximizer player's turn to choose an action, or -1 for the minimiser player -- `getPossibleActions()`: Returns an iterable of all actions which can be taken from this state +- `getPossibleActions()`: Returns an iterable of all `action`s which can be taken from this state - `takeAction(action)`: Returns the state which results from taking action `action` -- `isTerminal()`: Returns whether this state is a terminal state -- `getReward()`: Returns the reward for this state. Only needed for terminal states. +- `isTerminal()`: Returns `True` if this state is a terminal state +- `getReward()`: Returns the reward for this state. Only needed for terminal states. -You must also choose a hashable representation for an action as used in `getPossibleActions` and `takeAction`. Typically this would be a class with a custom `__hash__` method, but it could also simply be a tuple or a string. +You must also choose a hashable representation for an action as used in `getPossibleActions` and `takeAction`. Typically this would be a class with a custom `__hash__` method, but it could also simply be a tuple or a string. Once these have been implemented, running MCTS is as simple as initializing your starting state, then running: @@ -28,7 +28,7 @@ from mcts import mcts mcts = mcts(timeLimit=1000) bestAction = mcts.search(initialState=initialState) ``` -See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. +See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. ## Slow Usage //TODO From 7f76cb97b91d52268a765782166b1ea993c37c40 Mon Sep 17 00:00:00 2001 From: Lambda Date: Tue, 1 Dec 2020 13:53:28 +0800 Subject: [PATCH 02/12] add some README --- README.md | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index dc2ae78..1179a2d 100755 --- a/README.md +++ b/README.md @@ -25,9 +25,11 @@ Once these have been implemented, running MCTS is as simple as initializing your ```python from mcts import mcts -mcts = mcts(timeLimit=1000) -bestAction = mcts.search(initialState=initialState) +searcher = mcts(timeLimit=1000) +bestAction = searcher.search(initialState=initialState) ``` +Here the unit of `timeLimit=1000` is ms. You can also use `iterationLimit=1600` to specify number of roolouts. Only and at least one in `timeLimit` and `iterationLimit` should be specified. + See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. ## Slow Usage From c5854a2571ac0168c3a1d3a2bdd79cfb28f4bec5 Mon Sep 17 00:00:00 2001 From: WhymustIhaveaname Date: Tue, 1 Dec 2020 13:55:54 +0800 Subject: [PATCH 03/12] Update README.md --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 1179a2d..271a0f6 100755 --- a/README.md +++ b/README.md @@ -28,7 +28,7 @@ from mcts import mcts searcher = mcts(timeLimit=1000) bestAction = searcher.search(initialState=initialState) ``` -Here the unit of `timeLimit=1000` is ms. You can also use `iterationLimit=1600` to specify number of roolouts. Only and at least one in `timeLimit` and `iterationLimit` should be specified. +Here the unit of `timeLimit=1000` is millisecond. You can also use `iterationLimit=1600` to specify the number of rollouts. Only and at least one in `timeLimit` and `iterationLimit` should be specified. See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. From 7b80c8d0a3b2e2fc820b5a36d383624e6e1c4f7e Mon Sep 17 00:00:00 2001 From: Lambda Date: Fri, 4 Dec 2020 02:52:08 +0800 Subject: [PATCH 04/12] add needNodeValue feature --- mcts.py | 29 +++++++++++++++++++++-------- 1 file changed, 21 insertions(+), 8 deletions(-) mode change 100644 => 100755 mcts.py diff --git a/mcts.py b/mcts.py old mode 100644 new mode 100755 index 1db365a..714b5c9 --- a/mcts.py +++ b/mcts.py @@ -25,6 +25,13 @@ def __init__(self, state, parent): self.totalReward = 0 self.children = {} + def __str__(self): + s=[] + s.append("totalReward: %s"%(self.totalReward)) + s.append("numVisits: %d"%(self.numVisits)) + s.append("isTerminal: %s"%(self.isTerminal)) + s.append("children: %s"%(self.children.keys())) + return "%s: {%s}"%(self.__class__.__name__, ', '.join(s)) class mcts(): def __init__(self, timeLimit=None, iterationLimit=None, explorationConstant=1 / math.sqrt(2), @@ -46,7 +53,10 @@ def __init__(self, timeLimit=None, iterationLimit=None, explorationConstant=1 / self.explorationConstant = explorationConstant self.rollout = rolloutPolicy - def search(self, initialState): + def search(self, initialState, needNodeValue=False): + """ + do many executeRounds until the limit is reached, then return BestChild + """ self.root = treeNode(initialState, None) if self.limitType == 'time': @@ -58,9 +68,17 @@ def search(self, initialState): self.executeRound() bestChild = self.getBestChild(self.root, 0) - return self.getAction(self.root, bestChild) + action,node=((action, node) for action, node in self.root.children.items() if node is bestChild).__next__() + if needNodeValue: + nodeValue = node.totalReward / node.numVisits + return action, nodeValue + else: + return action def executeRound(self): + """ + execute a selection-expansion-simulation-backpropagation round + """ node = self.selectNode(self.root) reward = self.rollout(node.state) self.backpropogate(node, reward) @@ -102,9 +120,4 @@ def getBestChild(self, node, explorationValue): bestNodes = [child] elif nodeValue == bestValue: bestNodes.append(child) - return random.choice(bestNodes) - - def getAction(self, root, bestChild): - for action, node in root.children.items(): - if node is bestChild: - return action + return random.choice(bestNodes) \ No newline at end of file From a8732bd1e7e36c463b743b1307b666eace724c83 Mon Sep 17 00:00:00 2001 From: Lambda Date: Sat, 5 Dec 2020 01:36:49 +0800 Subject: [PATCH 05/12] modified NaughtsAndCrosses --- naughtsandcrosses.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) mode change 100644 => 100755 naughtsandcrosses.py diff --git a/naughtsandcrosses.py b/naughtsandcrosses.py old mode 100644 new mode 100755 index 5b4019a..9d490a3 --- a/naughtsandcrosses.py +++ b/naughtsandcrosses.py @@ -73,9 +73,9 @@ def __eq__(self, other): def __hash__(self): return hash((self.x, self.y, self.player)) +if __name__=="__main__": + initialState = NaughtsAndCrossesState() + searcher = mcts(timeLimit=1000) + action = searcher.search(initialState=initialState) -initialState = NaughtsAndCrossesState() -mcts = mcts(timeLimit=1000) -action = mcts.search(initialState=initialState) - -print(action) + print(action) From d911c49ef558c9d308aff42e871b346b76a03757 Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Sat, 9 Jan 2021 12:33:59 +0800 Subject: [PATCH 06/12] minor change as requested for pull 13 --- README.md | 7 ++++++- mcts.py | 14 +++++--------- 2 files changed, 11 insertions(+), 10 deletions(-) diff --git a/README.md b/README.md index 271a0f6..8436258 100755 --- a/README.md +++ b/README.md @@ -28,7 +28,12 @@ from mcts import mcts searcher = mcts(timeLimit=1000) bestAction = searcher.search(initialState=initialState) ``` -Here the unit of `timeLimit=1000` is millisecond. You can also use `iterationLimit=1600` to specify the number of rollouts. Only and at least one in `timeLimit` and `iterationLimit` should be specified. +Here the unit of `timeLimit=1000` is millisecond. You can also use `iterationLimit=1600` to specify the number of rollouts. Exactly one of `timeLimit` and `iterationLimit` should be specified. The expected reward of best action can be got by setting `needDetails` to `True` in `searcher`. + +```python +resultDict = searcher.search(initialState=initialState, needDetails=True) +print(resultDict.keys()) #currently includes dict_keys(['action', 'expectedReward']) +``` See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. diff --git a/mcts.py b/mcts.py index 714b5c9..3ea88f8 100755 --- a/mcts.py +++ b/mcts.py @@ -30,7 +30,7 @@ def __str__(self): s.append("totalReward: %s"%(self.totalReward)) s.append("numVisits: %d"%(self.numVisits)) s.append("isTerminal: %s"%(self.isTerminal)) - s.append("children: %s"%(self.children.keys())) + s.append("possibleActions: %s"%(self.children.keys())) return "%s: {%s}"%(self.__class__.__name__, ', '.join(s)) class mcts(): @@ -53,10 +53,7 @@ def __init__(self, timeLimit=None, iterationLimit=None, explorationConstant=1 / self.explorationConstant = explorationConstant self.rollout = rolloutPolicy - def search(self, initialState, needNodeValue=False): - """ - do many executeRounds until the limit is reached, then return BestChild - """ + def search(self, initialState, needDetails=False): self.root = treeNode(initialState, None) if self.limitType == 'time': @@ -68,10 +65,9 @@ def search(self, initialState, needNodeValue=False): self.executeRound() bestChild = self.getBestChild(self.root, 0) - action,node=((action, node) for action, node in self.root.children.items() if node is bestChild).__next__() - if needNodeValue: - nodeValue = node.totalReward / node.numVisits - return action, nodeValue + action=(action for action, node in self.root.children.items() if node is bestChild).__next__() + if needDetails: + return {"action": action, "expectedReward": bestChild.totalReward / bestChild.numVisits} else: return action From 3f15d1053ced598c79555ce066e7f781356dfdef Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Mon, 18 Jan 2021 23:51:40 +0800 Subject: [PATCH 07/12] update README: add write your own policy --- README.md | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/README.md b/README.md index 8436258..b08e90b 100755 --- a/README.md +++ b/README.md @@ -38,6 +38,30 @@ print(resultDict.keys()) #currently includes dict_keys(['action', 'expectedRewar See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. ## Slow Usage + +### Write Your Own Policy + +The default policy for this package is `randomPolicy` defined in `mcts.py`. Its structure is + +``` +def randomPolicy(state): + while not state.isTerminal(): + action = random.choice(state.getPossibleActions()) + state = state.takeAction(action) + return state.getReward() +``` + +By substituting it with a stronger policy, you can make the search more efficient. The new policy should be a function which takes `state` as its input and return reward from the point of view of `state`'s current player and will be hand over to mcts by changing `rolloutPolicy=randomPolicy` in `mcts`'s construct function. Pay attention to the sign of reward the policy function returned. Or it will play for its opponent. For example, suppose I have trained a neural network which can estimate the expected reward even the state is not terminal; I can use it to accelerate the rollout + +``` +def nnPolicy(state): + if state.isTerminal(): + return state.getReward() + else: + return reward_estimated_by_neural_network +``` + +### More //TODO ## Collaborating From 3ac90b946b4dd429f9d9a5603939b4259195d49d Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Wed, 31 Mar 2021 14:27:16 +0800 Subject: [PATCH 08/12] finish alphabeta --- mcts.py | 62 +++++++++++++++++++++++++++++++++++++++++++- naughtsandcrosses.py | 28 +++++++++++++++++--- 2 files changed, 85 insertions(+), 5 deletions(-) diff --git a/mcts.py b/mcts.py index 3ea88f8..0b9dcbe 100755 --- a/mcts.py +++ b/mcts.py @@ -116,4 +116,64 @@ def getBestChild(self, node, explorationValue): bestNodes = [child] elif nodeValue == bestValue: bestNodes.append(child) - return random.choice(bestNodes) \ No newline at end of file + return random.choice(bestNodes) + + +class abpruning(): + def __init__(self, deep=3, safemargin=0.1, gameinf=65535): + """ + deep: how many layers to be search, must >= 1 + safemargin: break condition beta <= alpha --> (beta + self.safemargin) <= alpha + gameinf: a number which will never be reached by getReward + """ + self.deep=deep + self.safemargin=safemargin + self.gameinf=gameinf + + def search(self, initialState, needDetails=False): + children={} + for action in initialState.getPossibleActions(): + val = self.alphabeta(initialState.takeAction(action), self.deep-1, -1*self.gameinf, self.gameinf) + children[action] = val + self.children = children + + """CurrentPlayer=initialState.getCurrentPlayer() + if CurrentPlayer==1: + bestaction = max(self.children.items(),key=lambda x: x[1]) + elif CurrentPlayer==-1: + bestaction = min(self.children.items(),key=lambda x: x[1]) + else: + raise Exception("getCurrentPlayer() should return 1 or -1 rather than %s"%(CurrentPlayer,)) + + if needDetails: + return {"action": bestaction[0], "expectedReward": bestaction[1]} + else: + return bestaction[0]""" + + def alphabeta(self, node, deep, alpha, beta): + if deep==0 or node.isTerminal(): + return node.getReward() + + CurrentPlayer=node.getCurrentPlayer() + if CurrentPlayer==1: + maxeval = -1*self.gameinf + actions = node.getPossibleActions() + for action in actions: + val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) + maxeval = max(val, maxeval) + alpha = max(val, alpha) + if (beta + self.safemargin) <= alpha: + break + return maxeval + elif CurrentPlayer==-1: + mineval = self.gameinf + actions = node.getPossibleActions() + for action in actions: + val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) + mineval = min(val, mineval) + beta = min(val, beta) + if (beta + self.safemargin) <= alpha: + break + return mineval + else: + raise Exception("getCurrentPlayer() should return 1 or -1 rather than %s"%(CurrentPlayer,)) \ No newline at end of file diff --git a/naughtsandcrosses.py b/naughtsandcrosses.py index 9d490a3..a626ed4 100755 --- a/naughtsandcrosses.py +++ b/naughtsandcrosses.py @@ -73,9 +73,29 @@ def __eq__(self, other): def __hash__(self): return hash((self.x, self.y, self.player)) -if __name__=="__main__": +def test_01(): initialState = NaughtsAndCrossesState() - searcher = mcts(timeLimit=1000) - action = searcher.search(initialState=initialState) - + # 1 first 1 step win + # deep=1,2: {(0, 1): False, (0, 2): 1.0, (1, 2): False, (2, 1): False, (2, 2): False} + # deep>=3 : {(0, 1): 1.0, (0, 2): 1.0, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} + #initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] + #initialState.currentPlayer = 1 + + # 1 first 3 step win + # deep>=3 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} + # deep=2 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): False, (2, 2): False} + initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] + initialState.currentPlayer = 1 + + from mcts import abpruning + searcher=abpruning(deep=2,safemargin=0.1,gameinf=65535) + action=searcher.search(initialState,needDetails=True) print(action) + print(searcher.children) + +if __name__=="__main__": + #initialState = NaughtsAndCrossesState() + #searcher = mcts(timeLimit=1000) + #action = searcher.search(initialState=initialState) + #print(action) + test_01() \ No newline at end of file From b13b133636389e1b432d17a6ebc04e6c005c211f Mon Sep 17 00:00:00 2001 From: WhymustIhaveaname Date: Wed, 31 Mar 2021 15:48:15 +0800 Subject: [PATCH 09/12] add alphabeta in README --- README.md | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/README.md b/README.md index b08e90b..9fe6bb2 100755 --- a/README.md +++ b/README.md @@ -37,6 +37,24 @@ print(resultDict.keys()) #currently includes dict_keys(['action', 'expectedRewar See [naughtsandcrosses.py](https://github.com/pbsinclair42/MCTS/blob/master/naughtsandcrosses.py) for a simple example. +### Alpha-Beta Pruning + +The use of alpha-beta pruning is almost the same as MCTS. The only different is that `getReward()` is needed for all states. + +```python +from mcts import abpruning +searcher=abpruning(deep=3) +bestAction=searcher.search(initialState) +``` + +The parameters for `abpruning`'s construction function are + +* deep : search deepth; +* safemargin: normally alpha-beta pruning will break when `beta <= alpha`, safemargin strengthen this to `beta + safemargin <= alpha` for situations where eval function is not very accurate; +* gameinf : an upper bound of getReward() return values used as "inf" in algorithm. + +Details of chlidren can be found in `searcher.children` after `search()` is called. `searcher.children` is a dictinary looks like {action:value}. + ## Slow Usage ### Write Your Own Policy From 6e40d497a82cd30a36c1a541ab47c1191cc4696f Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Wed, 31 Mar 2021 15:52:44 +0800 Subject: [PATCH 10/12] add counter --- chess.py | 2 ++ mcts.py | 13 +++++++++---- naughtsandcrosses.py | 12 ++++++------ 3 files changed, 17 insertions(+), 10 deletions(-) create mode 100644 chess.py diff --git a/chess.py b/chess.py new file mode 100644 index 0000000..0581e67 --- /dev/null +++ b/chess.py @@ -0,0 +1,2 @@ +class ChessState(): + pass \ No newline at end of file diff --git a/mcts.py b/mcts.py index 0b9dcbe..2cba90c 100755 --- a/mcts.py +++ b/mcts.py @@ -124,15 +124,17 @@ def __init__(self, deep=3, safemargin=0.1, gameinf=65535): """ deep: how many layers to be search, must >= 1 safemargin: break condition beta <= alpha --> (beta + self.safemargin) <= alpha - gameinf: a number which will never be reached by getReward + gameinf: an upper bound of getReward() return values used as "inf" in algorithm """ - self.deep=deep - self.safemargin=safemargin - self.gameinf=gameinf + self.deep = deep + self.safemargin = safemargin + self.gameinf = gameinf + self.counter = 0 def search(self, initialState, needDetails=False): children={} for action in initialState.getPossibleActions(): + self.counter += 1 val = self.alphabeta(initialState.takeAction(action), self.deep-1, -1*self.gameinf, self.gameinf) children[action] = val self.children = children @@ -152,6 +154,7 @@ def search(self, initialState, needDetails=False): def alphabeta(self, node, deep, alpha, beta): if deep==0 or node.isTerminal(): + self.counter += 1 return node.getReward() CurrentPlayer=node.getCurrentPlayer() @@ -159,6 +162,7 @@ def alphabeta(self, node, deep, alpha, beta): maxeval = -1*self.gameinf actions = node.getPossibleActions() for action in actions: + self.counter += 1 val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) maxeval = max(val, maxeval) alpha = max(val, alpha) @@ -169,6 +173,7 @@ def alphabeta(self, node, deep, alpha, beta): mineval = self.gameinf actions = node.getPossibleActions() for action in actions: + self.counter += 1 val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) mineval = min(val, mineval) beta = min(val, beta) diff --git a/naughtsandcrosses.py b/naughtsandcrosses.py index a626ed4..071a5ce 100755 --- a/naughtsandcrosses.py +++ b/naughtsandcrosses.py @@ -78,20 +78,20 @@ def test_01(): # 1 first 1 step win # deep=1,2: {(0, 1): False, (0, 2): 1.0, (1, 2): False, (2, 1): False, (2, 2): False} # deep>=3 : {(0, 1): 1.0, (0, 2): 1.0, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} - #initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] - #initialState.currentPlayer = 1 + initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] + initialState.currentPlayer = 1 # 1 first 3 step win # deep>=3 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} # deep=2 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): False, (2, 2): False} - initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] - initialState.currentPlayer = 1 + #initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] + #initialState.currentPlayer = 1 from mcts import abpruning - searcher=abpruning(deep=2,safemargin=0.1,gameinf=65535) + searcher=abpruning(deep=3,safemargin=0.1,gameinf=65535) action=searcher.search(initialState,needDetails=True) - print(action) print(searcher.children) + print(searcher.counter) if __name__=="__main__": #initialState = NaughtsAndCrossesState() From a45a49098ce7292cb69e8398072442229dd5b7d8 Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Thu, 1 Apr 2021 09:55:00 +0800 Subject: [PATCH 11/12] del safemargin --- README.md | 5 ++--- mcts.py | 20 +++++++++----------- naughtsandcrosses.py | 20 +++++++++++++------- 3 files changed, 24 insertions(+), 21 deletions(-) diff --git a/README.md b/README.md index 9fe6bb2..a716a64 100755 --- a/README.md +++ b/README.md @@ -50,10 +50,9 @@ bestAction=searcher.search(initialState) The parameters for `abpruning`'s construction function are * deep : search deepth; -* safemargin: normally alpha-beta pruning will break when `beta <= alpha`, safemargin strengthen this to `beta + safemargin <= alpha` for situations where eval function is not very accurate; -* gameinf : an upper bound of getReward() return values used as "inf" in algorithm. +* gameinf : an upper bound of getReward() return values, used as "inf" in algorithm. -Details of chlidren can be found in `searcher.children` after `search()` is called. `searcher.children` is a dictinary looks like {action:value}. +After `search()` is called, details of children can be found in `searcher.children`, and `searcher.counter` records how many leaf nodes are visited. `searcher.children` is a dictinary looks like {action:value}. ## Slow Usage diff --git a/mcts.py b/mcts.py index 2cba90c..b6d318f 100755 --- a/mcts.py +++ b/mcts.py @@ -118,23 +118,23 @@ def getBestChild(self, node, explorationValue): bestNodes.append(child) return random.choice(bestNodes) +def trivialPolicy(state): + return state.getReward() class abpruning(): - def __init__(self, deep=3, safemargin=0.1, gameinf=65535): + def __init__(self, deep=3, gameinf=65535, rolloutPolicy = trivialPolicy): """ deep: how many layers to be search, must >= 1 - safemargin: break condition beta <= alpha --> (beta + self.safemargin) <= alpha gameinf: an upper bound of getReward() return values used as "inf" in algorithm """ self.deep = deep - self.safemargin = safemargin + self.rollout = rolloutPolicy self.gameinf = gameinf self.counter = 0 def search(self, initialState, needDetails=False): children={} for action in initialState.getPossibleActions(): - self.counter += 1 val = self.alphabeta(initialState.takeAction(action), self.deep-1, -1*self.gameinf, self.gameinf) children[action] = val self.children = children @@ -155,29 +155,27 @@ def search(self, initialState, needDetails=False): def alphabeta(self, node, deep, alpha, beta): if deep==0 or node.isTerminal(): self.counter += 1 - return node.getReward() + return self.rollout(node) CurrentPlayer=node.getCurrentPlayer() - if CurrentPlayer==1: + if CurrentPlayer == 1: maxeval = -1*self.gameinf actions = node.getPossibleActions() for action in actions: - self.counter += 1 val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) maxeval = max(val, maxeval) alpha = max(val, alpha) - if (beta + self.safemargin) <= alpha: + if beta <= alpha: break return maxeval - elif CurrentPlayer==-1: + elif CurrentPlayer == -1: mineval = self.gameinf actions = node.getPossibleActions() for action in actions: - self.counter += 1 val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) mineval = min(val, mineval) beta = min(val, beta) - if (beta + self.safemargin) <= alpha: + if beta <= alpha: break return mineval else: diff --git a/naughtsandcrosses.py b/naughtsandcrosses.py index 071a5ce..efbf3e6 100755 --- a/naughtsandcrosses.py +++ b/naughtsandcrosses.py @@ -75,20 +75,26 @@ def __hash__(self): def test_01(): initialState = NaughtsAndCrossesState() - # 1 first 1 step win + initialState.currentPlayer = 1 + + # "1" first 1 step win # deep=1,2: {(0, 1): False, (0, 2): 1.0, (1, 2): False, (2, 1): False, (2, 2): False} # deep>=3 : {(0, 1): 1.0, (0, 2): 1.0, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} + # searcher.counter| no pruning| vanilla alphabeta + # deep=2 | 17 | 17 + # deep=3 | 49 | 31 initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] - initialState.currentPlayer = 1 - # 1 first 3 step win + # "1" first 3 step win + # deep=1,2: {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): False, (2, 2): False} # deep>=3 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} - # deep=2 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): False, (2, 2): False} - #initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] - #initialState.currentPlayer = 1 + # searcher.counter| no pruning| vanilla alphabeta + # deep=2 | 20 | 20 + # deep=3 | 60 | 43 + initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] from mcts import abpruning - searcher=abpruning(deep=3,safemargin=0.1,gameinf=65535) + searcher=abpruning(deep=3) action=searcher.search(initialState,needDetails=True) print(searcher.children) print(searcher.counter) From 2f14d90a9e3003dfab493a45c3d6c4c00aa5fe74 Mon Sep 17 00:00:00 2001 From: DELL-XPS Date: Thu, 1 Apr 2021 11:26:35 +0800 Subject: [PATCH 12/12] finish killer --- README.md | 5 +++-- mcts.py | 49 +++++++++++++++++++++++++++++--------------- naughtsandcrosses.py | 16 +++++++-------- 3 files changed, 44 insertions(+), 26 deletions(-) diff --git a/README.md b/README.md index a716a64..d9407bc 100755 --- a/README.md +++ b/README.md @@ -49,8 +49,9 @@ bestAction=searcher.search(initialState) The parameters for `abpruning`'s construction function are -* deep : search deepth; -* gameinf : an upper bound of getReward() return values, used as "inf" in algorithm. +* deep : search deepth; +* n_killer : number of killers in killer heuristic optimization, default is 2; +* gameinf : an upper bound of getReward() return values, used as "inf" in algorithm, default 65535. After `search()` is called, details of children can be found in `searcher.children`, and `searcher.counter` records how many leaf nodes are visited. `searcher.children` is a dictinary looks like {action:value}. diff --git a/mcts.py b/mcts.py index b6d318f..8fccadc 100755 --- a/mcts.py +++ b/mcts.py @@ -3,6 +3,7 @@ import time import math import random +import heapq def randomPolicy(state): @@ -122,61 +123,77 @@ def trivialPolicy(state): return state.getReward() class abpruning(): - def __init__(self, deep=3, gameinf=65535, rolloutPolicy = trivialPolicy): + def __init__(self, deep, rolloutPolicy = trivialPolicy, n_killer = 2, gameinf=65535): """ deep: how many layers to be search, must >= 1 gameinf: an upper bound of getReward() return values used as "inf" in algorithm """ self.deep = deep self.rollout = rolloutPolicy + self.n_killer = n_killer self.gameinf = gameinf self.counter = 0 def search(self, initialState, needDetails=False): - children={} + children = {} + killers = {} # best actions of brother branches, for killer heuristic optimization for action in initialState.getPossibleActions(): - val = self.alphabeta(initialState.takeAction(action), self.deep-1, -1*self.gameinf, self.gameinf) + val,ks = self.alphabeta(initialState.takeAction(action), self.deep-1, -1*self.gameinf, self.gameinf, killers = killers) children[action] = val + for k in ks: + killers[k] = killers.setdefault(k,0) + 1 self.children = children """CurrentPlayer=initialState.getCurrentPlayer() if CurrentPlayer==1: - bestaction = max(self.children.items(),key=lambda x: x[1]) + bestAction = max(self.children.items(),key=lambda x: x[1]) elif CurrentPlayer==-1: - bestaction = min(self.children.items(),key=lambda x: x[1]) + bestAction = min(self.children.items(),key=lambda x: x[1]) else: raise Exception("getCurrentPlayer() should return 1 or -1 rather than %s"%(CurrentPlayer,)) if needDetails: - return {"action": bestaction[0], "expectedReward": bestaction[1]} + return {"action": bestAction[0], "expectedReward": bestAction[1]} else: - return bestaction[0]""" + return bestAction[0]""" - def alphabeta(self, node, deep, alpha, beta): + def alphabeta(self, node, deep, alpha, beta, killers = {}): if deep==0 or node.isTerminal(): self.counter += 1 - return self.rollout(node) + return self.rollout(node),[] CurrentPlayer=node.getCurrentPlayer() + actions = node.getPossibleActions() + actions.sort(key=lambda x: killers.get(x,-1),reverse=True) + subkillers = {} + bestactions = [] if CurrentPlayer == 1: maxeval = -1*self.gameinf - actions = node.getPossibleActions() for action in actions: - val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) - maxeval = max(val, maxeval) + val,ks = self.alphabeta(node.takeAction(action), deep-1, alpha, beta, killers = subkillers) + maxeval = max(val,maxeval) alpha = max(val, alpha) + bestactions.append((action,val)) if beta <= alpha: break - return maxeval + for k in ks: + subkillers[k] = subkillers.setdefault(k,0) + 1 + bestactions.sort(key=lambda x: x[1],reverse=True) + bestactions = [i[0] for i in bestactions[0:min(len(bestactions),self.n_killer)]] + return maxeval,bestactions elif CurrentPlayer == -1: mineval = self.gameinf - actions = node.getPossibleActions() for action in actions: - val = self.alphabeta(node.takeAction(action), deep-1, alpha, beta) + val,ks = self.alphabeta(node.takeAction(action), deep-1, alpha, beta, killers = subkillers) mineval = min(val, mineval) beta = min(val, beta) + bestactions.append((action,val)) if beta <= alpha: break - return mineval + for k in ks: + subkillers[k] = subkillers.setdefault(k,0) + 1 + bestactions.sort(key=lambda x: x[1]) + bestactions = [i[0] for i in bestactions[0:min(len(bestactions),self.n_killer)]] + return mineval,bestactions else: raise Exception("getCurrentPlayer() should return 1 or -1 rather than %s"%(CurrentPlayer,)) \ No newline at end of file diff --git a/naughtsandcrosses.py b/naughtsandcrosses.py index efbf3e6..2339482 100755 --- a/naughtsandcrosses.py +++ b/naughtsandcrosses.py @@ -80,21 +80,21 @@ def test_01(): # "1" first 1 step win # deep=1,2: {(0, 1): False, (0, 2): 1.0, (1, 2): False, (2, 1): False, (2, 2): False} # deep>=3 : {(0, 1): 1.0, (0, 2): 1.0, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} - # searcher.counter| no pruning| vanilla alphabeta - # deep=2 | 17 | 17 - # deep=3 | 49 | 31 - initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] + # searcher.counter| no pruning| vanilla alphabeta| n_killer=2 + # deep=2 | 17 | 17 | 17 + # deep=3 | 49 | 31 | 28 + #initialState.board = [[-1, 0, 0], [-1, 1, 0], [1, 0, 0]] # "1" first 3 step win # deep=1,2: {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): False, (2, 2): False} # deep>=3 : {(0, 0): False, (0, 1): False, (1, 2): False, (2, 1): 1.0, (2, 2): 1.0} - # searcher.counter| no pruning| vanilla alphabeta - # deep=2 | 20 | 20 - # deep=3 | 60 | 43 + # searcher.counter| no pruning| vanilla alphabeta| n_killer=2 + # deep=2 | 20 | 20 | 20 + # deep=3 | 60 | 43 | 37 initialState.board = [[0, 0, -1], [-1, 1, 0], [1, 0, 0]] from mcts import abpruning - searcher=abpruning(deep=3) + searcher=abpruning(deep=3,n_killer=2) action=searcher.search(initialState,needDetails=True) print(searcher.children) print(searcher.counter)