Fixed a bug in the assignment of invalid turns. Added lots of documentation.
This commit is contained in:
1 file changed
+320
-90
+320
-90
@@ -95,6 +95,7 @@
|
|||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
|
"from multiprocessing import Pool\n",
|
||||||
"\n",
|
"\n",
|
||||||
"%load_ext blackcellmagic"
|
"%load_ext blackcellmagic"
|
||||||
]
|
]
|
||||||
@@ -152,7 +153,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 22,
|
"execution_count": 3,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
@@ -162,7 +163,8 @@
|
|||||||
"EXAMPLE_STACK_SIZE: Final[int] = 1000 # defines the game stack size for examples\n",
|
"EXAMPLE_STACK_SIZE: Final[int] = 1000 # defines the game stack size for examples\n",
|
||||||
"IMPOSSIBLE: Final[np.ndarray] = np.array([-1, -1], dtype=int)\n",
|
"IMPOSSIBLE: Final[np.ndarray] = np.array([-1, -1], dtype=int)\n",
|
||||||
"IMPOSSIBLE.setflags(write=False)\n",
|
"IMPOSSIBLE.setflags(write=False)\n",
|
||||||
"SIMULATE_TURNS: Final[int] = 70"
|
"SIMULATE_TURNS: Final[int] = 70\n",
|
||||||
|
"VERIFY_POLICY: Final[bool] = True"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -454,22 +456,22 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 23,
|
"execution_count": 11,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [
|
"outputs": [
|
||||||
{
|
{
|
||||||
"name": "stdout",
|
"name": "stdout",
|
||||||
"output_type": "stream",
|
"output_type": "stream",
|
||||||
"text": [
|
"text": [
|
||||||
"9.43 ms ± 1 ms per loop (mean ± std. dev. of 7 runs, 100 loops each)\n",
|
"9.31 ms ± 1.67 ms per loop (mean ± std. dev. of 7 runs, 100 loops each)\n",
|
||||||
"1 s ± 179 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
|
"831 ms ± 25.2 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"data": {
|
"data": {
|
||||||
"text/plain": "array([[[False, False, False, False, False, False, False, False],\n [False, False, False, False, False, False, False, False],\n [False, False, False, True, False, False, False, False],\n [False, False, True, False, False, False, False, False],\n [False, False, False, False, False, True, False, False],\n [False, False, False, False, True, False, False, False],\n [False, False, False, False, False, False, False, False],\n [False, False, False, False, False, False, False, False]]])"
|
"text/plain": "array([[[False, False, False, False, False, False, False, False],\n [False, False, False, False, False, False, False, False],\n [False, False, False, True, False, False, False, False],\n [False, False, True, False, False, False, False, False],\n [False, False, False, False, False, True, False, False],\n [False, False, False, False, True, False, False, False],\n [False, False, False, False, False, False, False, False],\n [False, False, False, False, False, False, False, False]]])"
|
||||||
},
|
},
|
||||||
"execution_count": 23,
|
"execution_count": 11,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"output_type": "execute_result"
|
"output_type": "execute_result"
|
||||||
}
|
}
|
||||||
@@ -558,7 +560,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": 12,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
@@ -593,7 +595,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": 13,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
@@ -664,15 +666,15 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 24,
|
"execution_count": 14,
|
||||||
"outputs": [
|
"outputs": [
|
||||||
{
|
{
|
||||||
"name": "stdout",
|
"name": "stdout",
|
||||||
"output_type": "stream",
|
"output_type": "stream",
|
||||||
"text": [
|
"text": [
|
||||||
"177 µs ± 3.97 µs per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n",
|
"172 µs ± 7.68 µs per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n",
|
||||||
"29.7 µs ± 106 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n",
|
"29.9 µs ± 1.08 µs per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n",
|
||||||
"31.2 µs ± 269 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n"
|
"31.6 µs ± 1.01 µs per loop (mean ± std. dev. of 7 runs, 10,000 loops each)\n"
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
@@ -748,7 +750,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": 15,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
@@ -760,13 +762,13 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": 28,
|
"execution_count": 16,
|
||||||
"outputs": [
|
"outputs": [
|
||||||
{
|
{
|
||||||
"name": "stdout",
|
"name": "stdout",
|
||||||
"output_type": "stream",
|
"output_type": "stream",
|
||||||
"text": [
|
"text": [
|
||||||
"95.1 ms ± 3.5 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)\n"
|
"89.4 ms ± 3.1 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)\n"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -874,7 +876,7 @@
|
|||||||
"## An abstract reversi game policy\n",
|
"## An abstract reversi game policy\n",
|
||||||
"\n",
|
"\n",
|
||||||
"For an easy use of policies an abstract class containing the policy generation / requests an action in an inherited instance of this class.\n",
|
"For an easy use of policies an abstract class containing the policy generation / requests an action in an inherited instance of this class.\n",
|
||||||
"This class filters the policy to only propose valid actions. Inherited instance do not need to care about this."
|
"This class filters the policy to only propose valid actions. Inherited instance do not need to care about this. This super class also manges exploration and exploitation with the epsilon value."
|
||||||
],
|
],
|
||||||
"metadata": {
|
"metadata": {
|
||||||
"collapsed": false
|
"collapsed": false
|
||||||
@@ -882,7 +884,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": 17,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
@@ -891,6 +893,20 @@
|
|||||||
" A game policy. Proposes where to place a stone next.\n",
|
" A game policy. Proposes where to place a stone next.\n",
|
||||||
" \"\"\"\n",
|
" \"\"\"\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
" def __init__(self, epsilon: float):\n",
|
||||||
|
" \"\"\"\n",
|
||||||
|
"\n",
|
||||||
|
" Args:\n",
|
||||||
|
" epsilon: the epsilon / greedy value. Should be between zero and one. Set the mixture of policy and exploration. One means only the policy is used. Zero means only random policies are used. All mixtures inbetween between are possible.\n",
|
||||||
|
" \"\"\"\n",
|
||||||
|
" if 0 > epsilon > 1:\n",
|
||||||
|
" raise ValueError(\"Epsilon should be between zero and one.\")\n",
|
||||||
|
" self._epsilon: float = epsilon\n",
|
||||||
|
"\n",
|
||||||
|
" @property\n",
|
||||||
|
" def epsilon(self):\n",
|
||||||
|
" return self._epsilon\n",
|
||||||
|
"\n",
|
||||||
" @property\n",
|
" @property\n",
|
||||||
" @abc.abstractmethod\n",
|
" @abc.abstractmethod\n",
|
||||||
" def policy_name(self) -> str:\n",
|
" def policy_name(self) -> str:\n",
|
||||||
@@ -909,39 +925,179 @@
|
|||||||
" \"\"\"\n",
|
" \"\"\"\n",
|
||||||
" raise NotImplementedError()\n",
|
" raise NotImplementedError()\n",
|
||||||
"\n",
|
"\n",
|
||||||
" def get_policy(\n",
|
" def get_policy(self, boards: np.ndarray) -> np.ndarray:\n",
|
||||||
" self, boards: np.ndarray, epsilon: float = 1\n",
|
" \"\"\"Calculates the policy that should be followed.\n",
|
||||||
" ) -> tuple[np.ndarray, np.ndarray]:\n",
|
"\n",
|
||||||
|
" Calculates the policy that should be followed.\n",
|
||||||
|
" This function does include the usage of epsilon to configure greediness and exploration.\n",
|
||||||
|
"\n",
|
||||||
|
" Args:\n",
|
||||||
|
" boards: A set of boards that show the environment where the policy should be calculated for.\n",
|
||||||
|
"\n",
|
||||||
|
" Returns:\n",
|
||||||
|
" A vector of indices. Should be formatted as an array of the form [x, y]. The value [-1, -1] is expected if no turn is possible.\n",
|
||||||
|
" \"\"\"\n",
|
||||||
" assert len(boards.shape) == 3\n",
|
" assert len(boards.shape) == 3\n",
|
||||||
" assert boards.shape == (BOARD_SIZE, BOARD_SIZE)\n",
|
" assert boards.shape[1:] == (BOARD_SIZE, BOARD_SIZE)\n",
|
||||||
"\n",
|
"\n",
|
||||||
" # todo possibly change this function to only validate the purpose turn and\n",
|
" if self.epsilon <= 0:\n",
|
||||||
|
" policies = np.random.rand(*boards.shape)\n",
|
||||||
|
" else:\n",
|
||||||
|
" policies = self._internal_policy(boards)\n",
|
||||||
|
" if self.epsilon < 1:\n",
|
||||||
|
" policies = policies * self.epsilon + np.random.rand(*boards.shape) * (\n",
|
||||||
|
" 1 - self.epsilon\n",
|
||||||
|
" )\n",
|
||||||
"\n",
|
"\n",
|
||||||
" policies = self._internal_policy(boards)\n",
|
" # todo talk to team about backpropagation of score and epsilon for greedy factor\n",
|
||||||
" raw_policy = policies.copy()\n",
|
|
||||||
" if epsilon < 1:\n",
|
|
||||||
" policies = policies + np.random.rand(*boards.shape)\n",
|
|
||||||
"\n",
|
|
||||||
" # todo talk to team about backpropagation epsilon for greedy factor\n",
|
|
||||||
"\n",
|
"\n",
|
||||||
|
" # todo possibly change this function to only validate the purpose turn and not all turns\n",
|
||||||
" possible_turns = get_possible_turns(boards)\n",
|
" possible_turns = get_possible_turns(boards)\n",
|
||||||
" policies[possible_turns == False] = -1.0\n",
|
" policies[possible_turns == False] = -1.0\n",
|
||||||
" max_indices = [\n",
|
" max_indices = [\n",
|
||||||
" np.unravel_index(policy.argmax(), policy.shape) for policy in policies\n",
|
" np.unravel_index(policy.argmax(), policy.shape) for policy in policies\n",
|
||||||
" ]\n",
|
" ]\n",
|
||||||
" policy_vector = np.array(max_indices)\n",
|
" policy_vector = np.array(max_indices)\n",
|
||||||
" max_policy = policy_vector\n",
|
" no_turn_possible_1 = np.all(policy_vector == 0, 1)\n",
|
||||||
|
" zero_pos = policies[:, 0, 0] == -1.0\n",
|
||||||
" no_turn_possible = np.all(policy_vector == 0, 1) & (policies[:, 0, 0] == -1.0)\n",
|
" no_turn_possible = np.all(policy_vector == 0, 1) & (policies[:, 0, 0] == -1.0)\n",
|
||||||
"\n",
|
"\n",
|
||||||
" policy_vector[no_turn_possible] = IMPOSSIBLE\n",
|
" policy_vector[no_turn_possible, :] = IMPOSSIBLE\n",
|
||||||
" max_policy[no_turn_possible] = 0\n",
|
" return policy_vector"
|
||||||
" return policy_vector, raw_policy"
|
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "markdown",
|
"cell_type": "markdown",
|
||||||
"source": [
|
"source": [
|
||||||
"## A first policy"
|
"## A first policy\n",
|
||||||
|
"\n",
|
||||||
|
"To quantify the quality of a game AI there needs to be some benchmarks.\n",
|
||||||
|
"The easiest benchmark is to play against a random player.\n",
|
||||||
|
"The easiest player to use as a benchmark is the random player.\n",
|
||||||
|
"For this and testing purpose the random policy was implemented."
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": 18,
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"class RandomPolicy(GamePolicy):\n",
|
||||||
|
" \"\"\"\n",
|
||||||
|
" A policy playing a random turn by setting epsilon to 0.\n",
|
||||||
|
" \"\"\"\n",
|
||||||
|
"\n",
|
||||||
|
" def __init__(self, epsilon: float):\n",
|
||||||
|
" _ = epsilon\n",
|
||||||
|
" super().__init__(epsilon=0)\n",
|
||||||
|
"\n",
|
||||||
|
" @property\n",
|
||||||
|
" def policy_name(self) -> str:\n",
|
||||||
|
" return \"random\"\n",
|
||||||
|
"\n",
|
||||||
|
" def _internal_policy(self, boards: np.ndarray) -> np.ndarray:\n",
|
||||||
|
" pass\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"rnd_policy = RandomPolicy(1)\n",
|
||||||
|
"assert rnd_policy.policy_name == \"random\"\n",
|
||||||
|
"assert rnd_policy.epsilon == 0\n",
|
||||||
|
"\n",
|
||||||
|
"rnd_policy_result = rnd_policy.get_policy(get_new_games(10))\n",
|
||||||
|
"assert np.any((5 >= rnd_policy_result) & (rnd_policy_result >= 3))"
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"source": [
|
||||||
|
"## Putting the game simulation together\n",
|
||||||
|
"Now it's time to bring all together for a proper simulation."
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"source": [
|
||||||
|
"### Playing a single turn\n",
|
||||||
|
"\n",
|
||||||
|
"The next function needed is used to request a policy, verify that the turn is legit and place a stone and turn enemy stones if possible."
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": 19,
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "stdout",
|
||||||
|
"output_type": "stream",
|
||||||
|
"text": [
|
||||||
|
"1.02 s ± 58.8 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n",
|
||||||
|
"949 ms ± 43.3 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)\n"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"data": {
|
||||||
|
"text/plain": "<Figure size 1200x600 with 8 Axes>",
|
||||||
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAABKUAAAJOCAYAAABm7rQwAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjYuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/P9b71AAAACXBIWXMAAA9hAAAPYQGoP6dpAABhL0lEQVR4nO3dfZCddX03/vfZLKxAsisgSGISBKGhAmEUtSgjgoo1kogdb9o68RawvX+txqcCtqYzWq2V6AiMvaumrUVCB/CpU6ziDSoqwQ7yqDzYakxqMYsJxWlxlwRdye75/XHM05KQPWd3r+u7Oa/XzBl3s+fs522uPW9OPnudcxrNZrMZAAAAAKhQT90BAAAAAOg+llIAAAAAVM5SCgAAAIDKWUoBAAAAUDlLKQAAAAAqZykFAAAAQOUspQAAAAConKUUAAAAAJXrrXrg2NhYNm3alDlz5qTRaFQ9HihQs9nMY489lnnz5qWnp75duX4C9qSEjtJPwJ7oJ6BUE+2nypdSmzZtyoIFC6oeC8wAg4ODmT9/fm3z9RPwVOrsKP0EPBX9BJRqX/1U+VJqzpw5Oz4+eG7V05PHH07STNJIDj6q+vkyyFBahrrnJ8njm1v/u2s/1KHufkoKOR5+JmWQYfcMBXSUfpKhlPkyFJZBPyUp5FjIIEMh84vJMMF+qnwptf2UzoPnJm/cVPX05Nr5ydafJofMS5Y/VP18GWQoLUPd85Pkmnmt0qr7lO+6+ykp43jUnaHu+TLIMF4JHaWfZChlvgxlZdBPLSUcCxlkKGV+KRkm2k9e6BwAAACAyllKAQAAAFA5SykAAAAAKmcpBQAAAEDlLKUAAAAAqJylFAAAAACVs5QCAAAAoHKWUgAAAABUru2l1K233pply5Zl3rx5aTQa+eIXvzgNsQDap5+AUuknoFT6CahT20uprVu35pRTTsknPvGJ6cgD0DH9BJRKPwGl0k9AnXrbvcGSJUuyZMmS6cgCMCn6CSiVfgJKpZ+AOnlNKQAAAAAq1/aZUu0aGRnJyMjIjs+Hh4eneyTAhOgnoFT6CSiVfgKm0rSfKbVq1aoMDAzsuCxYsGC6RwJMiH4CSqWfgFLpJ2AqTftSauXKlRkaGtpxGRwcnO6RABOin4BS6SegVPoJmErT/vS9vr6+9PX1TfcYgLbpJ6BU+gkolX4CplLbS6ktW7Zkw4YNOz7/z//8z9x777057LDDsnDhwikNB9AO/QSUSj8BpdJPQJ3aXkrdfffdOeuss3Z8ftFFFyVJzj///KxZs2bKggG0Sz8BpdJPQKn0E1CntpdSZ555ZprN5nRkAZgU/QSUSj8BpdJPQJ2m/YXOAQAAAGA8SykAAAAAKmcpBQAAAEDlLKUAAAAAqJylFAAAAACVs5QCAAAAoHKWUgAAAABUzlIKAAAAgMpZSgEAAABQuUaz2WxWOXB4eDgDAwNJIzlkXpWTWx7fnDTHkkZPcvDc6ufLIENpGeqenyRbNyVpJkNDQ+nv768nROrvp6SM41F3hrrnyyDDeCV0lH6SoZT5MpSVQT+1lHAsZJChlPmlZJhoP9W3lAIYp5ilFMAeFPGPPoA90E9AqfbVT70VZtmdM6VkkKGIDHXPT3Zu0YvhN31d/zMpgwy7Kqqj9FPXZ6h7vgxlZdBPLSUcCxlkKGV+KRkm2k+1LaUOPipZ/lD1c6+dn2z9aevA1DFfBhlKy1D3/CS5Zl6rOEtRVz8lZRyPujPUPV8GGcYrqaP0kwx1z5ehrAz6qaWEYyGDDKXMLyXDRPvJC50DAAAAUDlLKQAAAAAqZykFAAAAQOUspQAAAAConKUUAAAAAJWzlAIAAACgcpZSAAAAAFTOUgoAAACAyrW1lFq1alVe+MIXZs6cOTnyyCPzute9LuvWrZuubAATpp+AkukooFT6CahTW0uptWvXZsWKFbn99tvz9a9/PU888URe9apXZevWrdOVD2BC9BNQMh0FlEo/AXXqbefKN910026fr1mzJkceeWTuueeenHHGGVMaDKAd+gkomY4CSqWfgDpN6jWlhoaGkiSHHXbYlIQBmCr6CSiZjgJKpZ+AKrV1ptSuxsbG8q53vSunn356TjrppL1eb2RkJCMjIzs+Hx4e7nQkwIToJ6BkE+ko/QTUQT8BVev4TKkVK1bk+9//fj772c8+5fVWrVqVgYGBHZcFCxZ0OhJgQvQTULKJdJR+Auqgn4CqdbSUetvb3pYbbrgh3/rWtzJ//vynvO7KlSszNDS04zI4ONhRUICJ0E9AySbaUfoJqJp+AurQ1tP3ms1m3v72t+f666/PLbfckmOOOWaft+nr60tfX1/HAQEmQj8BJWu3o/QTUBX9BNSpraXUihUrct111+Vf/uVfMmfOnDz88MNJkoGBgRx00EHTEhBgIvQTUDIdBZRKPwF1auvpe6tXr87Q0FDOPPPMzJ07d8flc5/73HTlA5gQ/QSUTEcBpdJPQJ3afvoeQIn0E1AyHQWUSj8Bder43fcAAAAAoFOWUgAAAABUzlIKAAAAgMpZSgEAAABQOUspAAAAACpnKQUAAABA5SylAAAAAKicpRQAAAAAlWs0m81mlQOHh4czMDCQNJJD5lU5ueXxzUlzLGn0JAfPrX6+DDKUlqHu+UmydVOSZjI0NJT+/v56QqT+fkrKOB51Z6h7vgwyjFdCR+knGUqZL0NZGfRTSwnHQgYZSplfSoaJ9lN9SymAcYpZSgHsQRH/6APYA/0ElGpf/dRbYZbdOVNKBhmKyFD3/GTnFr0YftPX9T+TMsiwq6I6Sj91fYa658tQVgb91FLCsZBBhlLml5Jhov1U21Lq4KOS5Q9VP/fa+cnWn7YOTB3zZZChtAx1z0+Sa+a1irMUdfVTUsbxqDtD3fNlkGG8kjpKP8lQ93wZysqgn1pKOBYyyFDK/FIyTLSfvNA5AAAAAJWzlAIAAACgcpZSAAAAAFTOUgoAAACAytX37nsAANAltmxM1q1JhtYnTzyWHDAnGTg+WXRBMnth3ekAoB6WUgAAME02rU3uvzzZeEPrrbmTpDmaNGa1Pr7n/cnRS5PFlyRzz6gtJgDUwtP3AABgijWbyX2XJTecmQzemKTZWkY1R3/99e0fN5ONNyZffllredVs1hgaACpmKQUAAFPsgSuSO97d+ri57amvu/3rt1/Suh0AdAtLKQAAmEKb1rYWTJ24/ZJk861TmwcAStXWUmr16tVZvHhx+vv709/fnxe/+MW58cYbpysbwITpJ6BkOqq73H950ujwlVsbva3bQ1X0E1CntpZS8+fPz4c//OHcc889ufvuu/Pyl7885557bv7t3/5tuvIBTIh+Akqmo7rHlo2tFzXf11P29qa5LfnJl5Mtg1ObC/ZGPwF1amsptWzZsrzmNa/J8ccfn9/4jd/Ihz70ocyePTu33377dOUDmBD9BJRMR3WPdWt2vstepxo9ybqrpiQO7JN+AurU4YnFyejoaL7whS9k69atefGLX7zX642MjGRkZGTH58PDw52OBJgQ/QSUbCIdpZ9mrqH1U/N9hjdMzfeBdugnoGpt/x7ngQceyOzZs9PX15c//uM/zvXXX5/nPve5e73+qlWrMjAwsOOyYMGCSQUG2Bv9BJSsnY7STzPXE48lzdHJfY/maPIr/86nQvoJqEvbS6lFixbl3nvvzR133JG3vOUtOf/88/Pv//7ve73+ypUrMzQ0tOMyOOgJ8sD00E9AydrpKP00cx0wJ2nMmtz3aMxKDuyfmjwwEfoJqEvbT9878MADc9xxxyVJTj311Nx1113567/+6/zd3/3dHq/f19eXvr6+yaUEmAD9BJSsnY7STzPXwPFT8336j5ua7wMToZ+AukzyZRiTsbGx3Z5TDFAK/QSUTEftnxZdkDTHJvc9mmPJogunJA50RD8BVWnrTKmVK1dmyZIlWbhwYR577LFcd911ueWWW/LVr351uvIBTIh+Akqmo7rH7IXJwqXJ4I1Jc1v7t2/0Jgtfk8z2Mj1URD8BdWprKfXII4/kTW96UzZv3pyBgYEsXrw4X/3qV3P22WdPVz6ACdFPQMl0VHc55ZJk45c7u21zNFl88dTmgaein4A6tbWUuvLKK6crB8Ck6CegZDqqu8w9IzntsuT2S9q/7Wkfbd0eqqKfgDpN+jWlAACA3Z18UWsxlbSekvdUtn/9tMtatwOAbmEpBQAAU6zRaD0Nb9na1mtEpZE0ZrUuyS4fN1pfX7a2df1Go87UAFCttp6+BwAATNzcM1qXLYPJuquS4Q3Jr4aTA/uT/uNa77LnRc0B6FaWUgAAMM1mL0hOfV/dKQCgLJ6+BwAAAEDlLKUAAAAAqJylFAAAAACVs5QCAAAAoHKNZrPZrHLg8PBwBgYGkkZyyLwqJ7c8vjlpjiWNnuTgudXPl0GG0jLUPT9Jtm5K0kyGhobS399fT4jU309JGcej7gx1z5dBhvFK6Cj9JEMp82UoK4N+ainhWMggQynzS8kw0X6qbykFME4xSymAPSjiH30Ae6CfgFLtq596K8yyO2dKySBDERnqnp/s3KIXw2/6uv5nUgYZdlVUR+mnrs9Q93wZysqgn1pKOBYyyFDK/FIyTLSfaltKHXxUsvyh6udeOz/Z+tPWgaljvgwylJah7vlJcs28VnGWoq5+Sso4HnVnqHu+DDKMV1JH6ScZ6p4vQ1kZ9FNLCcdCBhlKmV9Khon2kxc6BwAAAKByllIAAAAAVM5SCgAAAIDKWUoBAAAAULn63n2PjmzZmKxbkwytT554LDlgTjJwfLLogmT2wu7IUPd8oFyHZkFekgtyZI7P0zInv8xjeSTrc1vW5NEMVpJBRwF7UkI3lJABKI/HT9TJUmqG2LQ2uf/yZOMNrbd1TJLmaNKY1fr4nvcnRy9NFl+SzD1j/8xQ93ygXMfnjJydi3NylqaZsSRJT3oy9uuPl+b9uT9fzs25POvz7WnJoKOAPSmhG0rIAJTH4ydK4OlLine truncated
|
||||||
|
},
|
||||||
|
"metadata": {},
|
||||||
|
"output_type": "display_data"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source": [
|
||||||
|
"def single_turn(\n",
|
||||||
|
" current_boards: np, policy: GamePolicy\n",
|
||||||
|
") -> tuple[np.ndarray, np.ndarray]:\n",
|
||||||
|
" \"\"\"Execute a single turn on a board.\n",
|
||||||
|
"\n",
|
||||||
|
" Places a new stone on the board. Turns captured enemy stones.\n",
|
||||||
|
"\n",
|
||||||
|
" Args:\n",
|
||||||
|
" current_boards: The current board before the game.\n",
|
||||||
|
" policy: The game policy to be used.\n",
|
||||||
|
"\n",
|
||||||
|
" Returns:\n",
|
||||||
|
" The new game board and the policy vector containing the index of the action used.\n",
|
||||||
|
" \"\"\"\n",
|
||||||
|
" policy_results = policy.get_policy(current_boards)\n",
|
||||||
|
"\n",
|
||||||
|
" # if the constant VERIFY_POLICY is set to true the policy is verified. Should be good though.\n",
|
||||||
|
" # todo deactivate the policy verification after some testing.\n",
|
||||||
|
" if VERIFY_POLICY:\n",
|
||||||
|
" assert np.all(moves_possible(current_boards, policy_results)), (\n",
|
||||||
|
" current_boards[(moves_possible(current_boards, policy_results) == False)],\n",
|
||||||
|
" policy_results[(moves_possible(current_boards, policy_results) == False)],\n",
|
||||||
|
" np.where(moves_possible(current_boards, policy_results) == False),\n",
|
||||||
|
" )\n",
|
||||||
|
" return do_moves(current_boards, policy_results), policy_results\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"%timeit single_turn(get_new_games(EXAMPLE_STACK_SIZE), RandomPolicy(1))\n",
|
||||||
|
"VERIFY_POLICY = False # type: ignore\n",
|
||||||
|
"%timeit single_turn(get_new_games(EXAMPLE_STACK_SIZE), RandomPolicy(1))\n",
|
||||||
|
"VERIFY_POLICY = True # type: ignore\n",
|
||||||
|
"plot_othello_boards(\n",
|
||||||
|
" single_turn(get_new_games(EXAMPLE_STACK_SIZE), RandomPolicy(1))[0][:8]\n",
|
||||||
|
")"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"source": [
|
||||||
|
"### Simulate a stack of games\n",
|
||||||
|
"This function will simulate a stack of games and return an array of policies and histories."
|
||||||
],
|
],
|
||||||
"metadata": {
|
"metadata": {
|
||||||
"collapsed": false
|
"collapsed": false
|
||||||
@@ -951,63 +1107,58 @@
|
|||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"outputs": [],
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "stderr",
|
||||||
|
"output_type": "stream",
|
||||||
|
"text": [
|
||||||
|
"Exception in thread Thread-5 (_handle_workers):\n",
|
||||||
|
"Traceback (most recent call last):\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\threading.py\", line 1016, in _bootstrap_inner\n",
|
||||||
|
" self.run()\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\threading.py\", line 953, in run\n",
|
||||||
|
" self._target(*self._args, **self._kwargs)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\pool.py\", line 516, in _handle_workers\n",
|
||||||
|
" cls._maintain_pool(ctx, Process, processes, pool, inqueue,\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\pool.py\", line 340, in _maintain_pool\n",
|
||||||
|
" Pool._repopulate_pool_static(ctx, Process, processes, pool,\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\pool.py\", line 329, in _repopulate_pool_static\n",
|
||||||
|
" w.start()\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\process.py\", line 121, in start\n",
|
||||||
|
" self._popen = self._Popen(self)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\context.py\", line 336, in _Popen\n",
|
||||||
|
" return Popen(process_obj)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\popen_spawn_win32.py\", line 93, in __init__\n",
|
||||||
|
" reduction.dump(process_obj, to_child)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\reduction.py\", line 60, in dump\n",
|
||||||
|
" ForkingPickler(file, protocol).dump(obj)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\synchronize.py\", line 104, in __getstate__\n",
|
||||||
|
" h = context.get_spawning_popen().duplicate_for_child(sl.handle)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\popen_spawn_win32.py\", line 99, in duplicate_for_child\n",
|
||||||
|
" return reduction.duplicate(handle, self.sentinel)\n",
|
||||||
|
" File \"C:\\Program Files\\Python310\\lib\\multiprocessing\\reduction.py\", line 79, in duplicate\n",
|
||||||
|
" return _winapi.DuplicateHandle(\n",
|
||||||
|
"PermissionError: [WinError 5] Zugriff verweigert\n"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
"source": [
|
"source": [
|
||||||
"class RandomPolicy(GamePolicy):\n",
|
"from tqdm.notebook import tqdm\n",
|
||||||
" @property\n",
|
|
||||||
" def policy_name(self) -> str:\n",
|
|
||||||
" return \"random\"\n",
|
|
||||||
"\n",
|
|
||||||
" def internal_policy(self, boards: np.ndarray) -> np.ndarray:\n",
|
|
||||||
" random_values = np.random.rand(*boards.shape)\n",
|
|
||||||
" return random_values\n",
|
|
||||||
" # return np.argmax(random_values, (1, 2))\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"rnd_policy = RandomPolicy()\n",
|
|
||||||
"assert rnd_policy.policy_name == \"random\"\n",
|
|
||||||
"rnd_policy_result = rnd_policy.get_policy(get_new_games(1))\n",
|
|
||||||
"assert np.any((5 >= rnd_policy_result) & (rnd_policy_result >= 3))"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"def single_turn(\n",
|
|
||||||
" current_boards: np, policy: GamePolicy\n",
|
|
||||||
") -> tuple[np.ndarray, np.ndarray]:\n",
|
|
||||||
" policy_results = policy.get_policy(current_boards)\n",
|
|
||||||
"\n",
|
|
||||||
" assert np.all(moves_possible(current_boards, policy_results)), (\n",
|
|
||||||
" current_boards[(moves_possible(current_boards, policy_results) == False)],\n",
|
|
||||||
" policy_results[(moves_possible(current_boards, policy_results) == False)],\n",
|
|
||||||
" np.where(moves_possible(current_boards, policy_results) == False),\n",
|
|
||||||
" )\n",
|
|
||||||
"\n",
|
|
||||||
" return do_moves(current_boards, policy_results), policy_results\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"%timeit single_turn(get_new_games(100), RandomPolicy())\n",
|
|
||||||
"single_turn(get_new_games(100), RandomPolicy())[0]"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"\n",
|
|
||||||
"\n",
|
"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"def simulate_game(\n",
|
"def simulate_game(\n",
|
||||||
" nr_of_games: int,\n",
|
" nr_of_games: int,\n",
|
||||||
" policies: tuple[GamePolicy, GamePolicy],\n",
|
" policies: tuple[GamePolicy, GamePolicy],\n",
|
||||||
") -> tuple[np.ndarray, np.ndarray]:\n",
|
") -> tuple[np.ndarray, np.ndarray]:\n",
|
||||||
|
" \"\"\"Simulates a stack of games.\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
" Args:\n",
|
||||||
|
" nr_of_games: The number of games that should be simulated.\n",
|
||||||
|
" policies: The policies that should be used to simulate the game.\n",
|
||||||
|
"\n",
|
||||||
|
" Returns:\n",
|
||||||
|
" A stack of board histories and actions.\n",
|
||||||
|
" \"\"\"\n",
|
||||||
" board_history_stack = np.zeros((SIMULATE_TURNS, nr_of_games, 8, 8))\n",
|
" board_history_stack = np.zeros((SIMULATE_TURNS, nr_of_games, 8, 8))\n",
|
||||||
" action_history_stack = np.zeros((SIMULATE_TURNS, nr_of_games, 2))\n",
|
" action_history_stack = np.zeros((SIMULATE_TURNS, nr_of_games, 2))\n",
|
||||||
" current_boards = get_new_games(nr_of_games)\n",
|
" current_boards = get_new_games(nr_of_games)\n",
|
||||||
@@ -1026,21 +1177,88 @@
|
|||||||
" return board_history_stack, action_history_stack\n",
|
" return board_history_stack, action_history_stack\n",
|
||||||
"\n",
|
"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"%timeit simulate_game(100, (RandomPolicy(), RandomPolicy()))\n",
|
"simulation_results = simulate_game(1, (RandomPolicy(1), RandomPolicy(1)))"
|
||||||
"simulate_game(10, (RandomPolicy(), RandomPolicy()))"
|
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": []
|
"source": [
|
||||||
|
"\n",
|
||||||
|
"%timeit simulate_game(100, (RandomPolicy(1), RandomPolicy(1)))\n",
|
||||||
|
"# simulate_game(EXAMPLE_STACK_SIZE, (RandomPolicy(1), RandomPolicy(1)))"
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false
|
||||||
|
}
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"policies_to_use = RandomPolicy(1), RandomPolicy(1)\n",
|
||||||
|
"with Pool(3) as pool:\n",
|
||||||
|
" results = pool.map(simulate_game, [100, policies_to_use])"
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false,
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"a = np.array(\n",
|
||||||
|
" [\n",
|
||||||
|
" [\n",
|
||||||
|
" [-1, -1, -1, -1, 0, 0, 0, 0],\n",
|
||||||
|
" [1, 1, -1, 1, 1, 0, 0, 0],\n",
|
||||||
|
" [1, 1, -1, 1, 1, 1, 0, 0],\n",
|
||||||
|
" [0, 1, -1, 1, 1, 1, 0, 0],\n",
|
||||||
|
" [0, 1, 1, 1, 1, 1, 0, 0],\n",
|
||||||
|
" [-1, 1, 1, 1, 1, 0, 0, 0],\n",
|
||||||
|
" [0, 0, 0, 1, 0, 0, 0, 0],\n",
|
||||||
|
" [0, 0, 0, 0, 0, 0, 0, 0],\n",
|
||||||
|
" ]\n",
|
||||||
|
" ],\n",
|
||||||
|
" dtype=int,\n",
|
||||||
|
")\n",
|
||||||
|
"a"
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"collapsed": false,
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"RandomPolicy(1).get_policy(a)"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"import numpy as np\n",
|
"import numpy as np\n",
|
||||||
@@ -1291,7 +1509,11 @@
|
|||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
"metadata": {
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"plot_othello_boards(create_test_game()[-3:])"
|
"plot_othello_boards(create_test_game()[-3:])"
|
||||||
@@ -1300,7 +1522,11 @@
|
|||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
"metadata": {
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": [
|
"source": [
|
||||||
"array = create_test_game()"
|
"array = create_test_game()"
|
||||||
@@ -1321,7 +1547,11 @@
|
|||||||
{
|
{
|
||||||
"cell_type": "code",
|
"cell_type": "code",
|
||||||
"execution_count": null,
|
"execution_count": null,
|
||||||
"metadata": {},
|
"metadata": {
|
||||||
|
"pycharm": {
|
||||||
|
"is_executing": true
|
||||||
|
}
|
||||||
|
},
|
||||||
"outputs": [],
|
"outputs": [],
|
||||||
"source": []
|
"source": []
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in new issue
Block a user