diff --git a/CoordConv.ipynb b/CoordConv.ipynb index 1d6cc7e..eb47dce 100644 --- a/CoordConv.ipynb +++ b/CoordConv.ipynb @@ -175,7 +175,7 @@ { "data": { "text/plain": [ - "" + "" ] }, "execution_count": 3, @@ -529,16 +529,16 @@ "name": "stdout", "output_type": "stream", "text": [ - "Train Epoch: 1 [2508/2508 (100%)] Loss: 7.505956\n", - "Train Epoch: 2 [2508/2508 (100%)] Loss: 4.008380\n", - "Train Epoch: 3 [2508/2508 (100%)] Loss: 2.033982\n", - "Train Epoch: 4 [2508/2508 (100%)] Loss: 0.994723\n", - "Train Epoch: 5 [2508/2508 (100%)] Loss: 0.472303\n", - "Train Epoch: 6 [2508/2508 (100%)] Loss: 0.254768\n", - "Train Epoch: 7 [2508/2508 (100%)] Loss: 0.135046\n", - "Train Epoch: 8 [2508/2508 (100%)] Loss: 0.087108\n", - "Train Epoch: 9 [2508/2508 (100%)] Loss: 0.060861\n", - "Train Epoch: 10 [2508/2508 (100%)] Loss: 0.046419\n" + "Train Epoch: 1 [2508/2508 (100%)] Loss: 7.513226\n", + "Train Epoch: 2 [2508/2508 (100%)] Loss: 4.021207\n", + "Train Epoch: 3 [2508/2508 (100%)] Loss: 2.022669\n", + "Train Epoch: 4 [2508/2508 (100%)] Loss: 1.027538\n", + "Train Epoch: 5 [2508/2508 (100%)] Loss: 0.476256\n", + "Train Epoch: 6 [2508/2508 (100%)] Loss: 0.249001\n", + "Train Epoch: 7 [2508/2508 (100%)] Loss: 0.141679\n", + "Train Epoch: 8 [2508/2508 (100%)] Loss: 0.089806\n", + "Train Epoch: 9 [2508/2508 (100%)] Loss: 0.060917\n", + "Train Epoch: 10 [2508/2508 (100%)] Loss: 0.045695\n" ] } ], @@ -557,11 +557,14 @@ " net.eval()\n", " test_loss = 0\n", " correct = 0\n", + " pred_logits = torch.tensor([])\n", " for data, target in test_loader:\n", " with torch.no_grad():\n", " data, target = Variable(data), Variable(target)\n", " data, target = data.to(device), target.to(device)\n", " output = net(data)\n", + " logits = F.softmax(output, dim=1)\n", + " pred_logits = torch.cat((pred_logits, logits.cpu()), dim=0)\n", " test_loss += criterion(output, target).item()\n", " _, pred = output.max(1, keepdim=True)\n", " _, label = target.max(dim=1)\n", @@ -571,7 +574,10 @@ " test_loss /= len(test_loader) # loss function already averages over batch size\n", " print('\\nTest set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%)\\n'.format(\n", " test_loss, correct, len(test_loader.dataset),\n", - " 100. * correct / len(test_loader.dataset)))" + " 100. * correct / len(test_loader.dataset)))\n", + " pred_logits = torch.sum(pred_logits, dim=0)\n", + " plt.imshow(pred_logits.detach().numpy().reshape(-1,64), cmap='gray')\n", + " plt.title('Predictions One-hot dataset')" ] }, { @@ -584,9 +590,19 @@ "output_type": "stream", "text": [ "\n", - "Test set: Average loss: 0.0461, Accuracy: 628/628 (100%)\n", + "Test set: Average loss: 0.0459, Accuracy: 628/628 (100%)\n", "\n" ] + }, + { + "data": { + "image/png": "iVBORw0KGgoAAAANSUhEUgAAAP4AAAEICAYAAAB/KknhAAAABHNCSVQICAgIfAhkiAAAAAlwSFlzAAALEgAACxIB0t1+/AAAADl0RVh0U29mdHdhcmUAbWF0cGxvdGxpYiB2ZXJzaW9uIDIuMi4yLCBodHRwOi8vbWF0cGxvdGxpYi5vcmcvhp/UCwAAIABJREFUeJztnX3UZWV12H8bhpnhY2BmUFgDow6K31kJIhItNjEEjZo0srrUatJ26qKZpMsuYRkraNNGU1vRtEqbDw0Rk7HLL6IihFWjqFiTNEEHlQgiggRhwsio8wGDEPnY/eOcd9h3e599n3vec+/7vpz9W2vW3HvPuc+zz3PO8969n72fvUVVSZJkWByy1AIkSTJ/cuInyQDJiZ8kAyQnfpIMkJz4STJAcuInyQDJiT8FIrJFRFREVrXvPyUiWzu083gROSAih/Yv5fJBRG4TkbPm1NfIvUliHnUTv33Y7msn1l0i8icictQs+lLVl6jq9kqZDk4AVb1dVY9S1YdmIVcgxz8Rkc+LyD0isl9E/lxEnjFPGWppJ/HJM2r7BSKycxZtL0U/XXjUTfyWf6aqRwGnAs8BfsufIA2P1uv/MUTkecBngMuBE4CTgOuAvxaRJy6lbMkSoKqPqn/AbcBZ5v3vAle2r78A/Ffgr4H7gJOBY4BLgF3APwBvAw5tzz8U+O/A94FbgdcCCqwy7f1b09evATcC9wDfoPnD87+Bh9v+DgBvBLa4dk4ArgD2ALcAv2bafAtwKfCBtt0bgNPM8fNbue8BbgJ+vjAufwn84ZjPPwV8oH39AmAn8JvA7nZMXmPOXdOOx+3AXcB7gcMn3Is3AH8H7Ac+Cqx143VLe91XACe0n3+xHZ972zH7F2PannRvXmPuxa3Ar7efH9nei4fbtg+043868DfAvva6fx9Y3X5HgHe3Y7K/vZ6fiMak1M9Sz4+D47fUAsxy4gOPayfKfzET9XbgmcAq4DDgk8AftTfqOOBL5iH5DeCbbTsbgaspTHzgFe0EfE77oJwMPMHL1L7f4tr5v8AfAmuBU4Dv0U5gmol/P/DS9mF/O/C37bGnAneYCbMFeNKYMTkCeAj4uTHHXgPsMhP/QeB32rF5KfBDYEN7/CKaCboRWAf8OfD2CffiS+3E2kgzEX+jPXYmzaQ9tZ08vwd80XxXgZODtifdm18EntTei59tr+NUc507XXvPBp7bPhdbWlnPa4/9AnAtsL5t7+nApkljMq6f5fJvyQXo/YKah+0AzV/u77QT6vD22BeA3zHnHg/8I+ZXC3g1cHX7+vMLD2r7/kWUJ/6ngXMDmcZO/PbBfQhYZ46/HfjT9vVbgM+aY88A7mtfn0zzK3QWcFgwJpvb/p425tiLgQfMg3rfwvW1n+1uJ4TQ/AI/yRx7HvD3E+7FvzTv3wm8t319CfBOc+wo4AFgS/t+0sQP782Y8z+5cH9qJiRwHnBZ+/pM4FvtOBxizgnHpKafpfr3aF0BPVtVP1s4dod5/QSaX7ZdIrLw2SHmnBPc+d8J+nwc8O3pReUEYI+q3uP6Oc28/655/UNgrYisUtVbROQ8mj8OzxSRTwOvV9U7XR97aVTOTTS/kpZNNL+8C/xAVR90/R0FPJZGc7jWjJXQaCGIyKeAf9p+/uuq+sGC7CeY6/7KwgFVPSAiPwBOpPmDMYnw3ojIS4DfBp5Cc0+PAL5eakxEngK8i2bcj6D5o3xtK9vnReT3gT8AHi8il9GYMGsJxmQ5M5jFLYPdjngHzS/+Y1R1ffvvaFV9Znt8F82EXuDxQbt30KiWk/r03AlsFJF1rp9/CL7zSMOqH1LV59P8EVPgHWPOuZfGfn3FmCZeCXyuoqvv02gDzzRjdYw2i6ho4+E4qv33wbgpoLnuJyy8EZEjgWOpvG6CeyMia4CP09jex6vqeuD/0ExKGH8/3kPzR/HJqno08GZzPqr6v1T12TRm4lOA/8CEMSn0sywY4sQ/iKruolnp/h8icrSIHCIiTxKRn21PuRR4nYhsFpENwAVBc+8D3iAiz249BieLyMKDfRcwduVcVe8A/h/wdhFZKyI/CZwDTJw8IvJUETmzfdDvp3kISy7CC4CtIvI6EVknIhtE5G00qulbJ/Wlqg8Dfwy8W0SOa/s/UUR+YdJ3C3wIeI2InNLK/9+Aa1T1tvZ4ccxaonuzmmbd4HvAg+2v/4vM8buAY0XkGPPZOuBu4ICIPA34dwsHROQ5IvLTInIYjWp/P/BQxZiM62dZMOiJ3/KvaR6Ub9CoxB+jUX+huamfpnF7fQX4RKkRVf0zGo/Bh2hWkj9Js+ADjc3+WyKyT0TeMObrr6ax++8ELgN+W1WvqpB9DXAhzS/Pd2kWJ99ckO+vaBap/jnNr+V3gGcBz1fVmyv6gsaDcAvwtyJyN/BZmgXGqVHVzwH/ieaXeReNtvQqc8pbgO3tmL1yTBPFe9OaTa+j+eOwF/gVmgW4hePfBD4M3Nq2fwKN6v4rNPfuj2k8EAsc3X62l2bcfkCjTUAwJoV+lgXSLkIkSTIg8hc/SQZITvwkGSA58ZNkgCxq4ovIi0XkJhG5RUSiFe8kSZYRnRf32i2l3wJeSBPf/WXg1ar6jeA7uZKYJDNGVWXSOYv5xT8duEVVb1XVHwEfAV62iPaSJJkTi5n4JzIaMrmz/WwEEdkmIjtEZMci+kqSpEcWE6s/Tp34MVVeVS8GLoZU9ZNkubCYX/ydjMZKb6aJPEuSZJmzmIn/ZeDJInKSiKymCbe8YsJ3kiRZBnRW9VX1QRH59zTx0ocC71fVG3qTLEmSmTHXWP208ZNk9szanZckyQolJ36SDJAVl3rLpDhi1maK7cuziIjHRbdR276lj766jkff19zHffFtzPN+9tX3Yshf/CQZIDnxk2SA5MRPkgGyImz8Qw99JFvx3XffffD10UcfPXLeQw8tvhSd7euHP/xh8bwjjjii2K+14Wx7MCr/unXrRo7Vym/bP+SQ0b/dJZmtvNP01WU8PPfc80jm8OiaI9u3Vo4jjzyy2L4dKyuTlysam1WrRqdM6X76Nqz8Xfvuk/zFT5IBkhM/SQbIiovc6+oa6uLm8t8p9e3V+XvvvbfY5lFHPVKx+4EHHiieZ/HtR6rzww8/PFHeSZTOrR0P/z5yv9lrO3DgwMgxO1YPPvhIcZ9IDm/6lExDO05eXn8tkZp+zDGPpMy3Mk7zXPU9BzNyL0mSseTET5IB8qhV9b163Pdqt13dteokjKp/XqW077uOvVdnS+1b/Gq0NUe86mlXxmvV11p8X/Za7r///pFjtr/DDz/84OtpVr5LHhZ/z2pX1ruM/bxJVT9JkrHkxE+SAZITP0kGyIqw8a19un///oOvrS0N3SLfIjeU57DDDjv4OoogtG10tSUjGe14ePeSdYHZNvx5XmZLaT3E2tkwav9HRO4wK0cUARndly7PcFeXWq1rteu19EHa+EmSjCUnfpIMkGW5ScerU1a9X79+ffF73h1kiTbVlLCqPZRdYF51syqwVb09Xt0sqcS+jSgqzrqbSlF8/r2PILRjdd999419DbGLzcpoj3kTw16nP2bH0V6Xj/CzKrY3P0r3qau67a/Tuj5tm959amX297PWZOqT/MVPkgGSEz9JBkhO/CQZICvCnRft7qr9Tq19FyV8sHamte2iXXY+xNPayV5Ga6vu3bt37OcwahN6W7K0NhAl7PDjYW33aJ3AttnVbWnljxJsWJm8HPaa/RqQHcfIddh3sk2/TmXXn6Lx7iMRR7rzkiQZy8SJLyLvF5HdInK9+WyjiFwlIje3/2+YrZhJkvTJRFVfRH4GOAB8QFV/ov3sncAeVb1QRC4ANqjq+RM762F3XrTDykby1apykUlQm3jCq9FRogx7rneP2TYjU6JLDvsoWYiPgPzRj340dV+ePqLpal1xtaZgtKPSmkWRut1H5F4kbx87/HpR9VX1i8Ae9/HLgO3t6+3A2VNLlyTJktE1gOd4Vd0FoKq7ROS40okisg3Y1rGfJElmQNWqvohsAa40qv4+VV1vju9V1Yl2ft/VcqPIt8gM6GMTTYRV56dRlUubkXykVxe5aqMEfX+2r0jNjVKd1+bfi1T4PspT2fGNov8iL00XD5OXw0eEllJ0d43om+Wq/l0isgmg/X93x3aSJFkCuk78K4Ct7eutwOX9iJMkyTyoced9GPgb4KkislNEzgEuBF4oIjcDL2zfJ0myQlgRkXtdiHbW2WuOdkrV2pyR7bthw+jSh40ki47ZNYna/PueKGFHbY75KHKv1BeM2q02ms7b1rYv7/rct2/f2GPTlKeyfdt1Ez/2e/Y84riKdlRGawP2Wrwctj97nwE2btw49ljXtZ2M3EuSZCw58ZNkgCzLRBy1RK4brx6vXbt27PeinHu1bpcoOYPHqm++70jFrKUUneZVW2v6eDXdXo+VyauodgyiiDM73n6sSqqyJzJbrDofRf9F7t7ajTLeHCltivLPQCkq02OPzTJPf/7iJ8kAyYmfJAMkJ36SDJAVYeOX7LTanPL+XGsTervV2sLeDrShuLW57T3Rzj0rc21yxigRh3XZ+fBPu+YRtWHt3chu9ZSSUPrc/BG2DXvfpynXba979erVxTYi7Ll+HG1i0i67JqH8TPvno48kHQvkL36SDJCc+EkyQJalqh9Fwtkop2lUPnvMukki14pXS62qHyVWsPJ71bDWfWXbiPLZ15aWipKFeFdfbbSePVabjGSaqDsrl426i/Lv+7LnlijaMnJ92gjC2mi62oQdvo2SidQ3+YufJAMkJ36SDJAVsUnHqpH2tVcNoyQGpVX4SHXz6ppN/2zVXK96WrXRr5hbohViK0e0AcbLb9X0rvnyLFG6cYs3i2qjziJzpJSf0N8XO45ejtL9nKZ6sPWweDW9dgNVbVKRPqrq5iadJEnGkhM/SQZITvwkGSDL0saPXC21CSQivC1pqS0tXdtvFBVXmwwiytE+5/tXfB/Z9LUuO+s2g9ExqL0v0W7LKOlnNI52l2bXfPzzjNxLGz9JkrHkxE+SAbIsVX2PVYVK+eC6thepnl6ts+payfXmZfTUutHmeV88XUyarm1H96Lkno3aiFyO9p5FOeunqR9QaidywUYqvDVDuz7fqeonSTKWnPhJMkBy4ifJAFmWNn5tqKzfiRXZgbVlsi2R289+b5rQ4dq+LbO4R6XkJvDjuwEXsMk7YNQ2ncY9ViOTx8p4//33F/vyIbu2zSjUOXKV1drdtSHBfYSTR6SNnyTJWGpKaD1ORK4WkRtF5AYRObf9fKOIXCUiN7f/T6yWmyTJ8mCiqt9Ww92kql8RkXXAtcDZwL8B9qjqhSJyAbBBVc+f0FYnndWqa6W88RDn0rOqVqmEs2+/djdalN8vyr/nVWrbd5Rcokv55GnKfNmc+/baoh14fZQbj8yFSNW33/PmSMn86xptGanf05QAt9hri3YC1t73XlR9Vd2lql9pX98D3AicCLwM2N6etp3mj0GSJCuAqVJvicgW4FnANcDxqroLmj8OInJc4TvbgG2LEzNJkj6pnvgichTwceA8Vb17ihXpi4GL2zaWLhwtSZKDVE18ETmMZtJ/UFU/0X58l4hsan/tNwG7ZyWktY9qEyt6W8zabVHOetuXt01LOwO9/Vaba93Lb+WySUV9aWabkLHWfo7q+0UZbawcfs3DXpu/lpLbKwq3rV0nWLNmzdi2J1G7O2+akOsu9rlvw15b7bOzWGpW9QW4BLhRVd9lDl0BbG1fbwUu71+8JElmQc0v/hnAvwK+LiJfaz97M3AhcKmInAPcDrxiNiImSdI3yzJyz9Mlssm7ZKyLyiZWiBJeRlFatTunpnHxlNyWXtW31x0le+yasKMU1Re50bz6bcc12j1no+J8jYOSejxNMg+vci/gx7Q22rKPUt6RjF1ctZ6M3EuSZCw58ZNkgKwIVd+qilZl92qcVb+9qnXMMcccfB2thNeWnaotC+VV29pNRpGKbfGRajXyjpO5RLSZx94Lb3aVTBBv+lj1uLQ5CEavMzKzontm1W2v6kdlz2zEYtcoREsUGdjFY+NJVT9JkrHkxE+SAZITP0kGyIqw8S2Ry6SUhz2ia2LM2nGLbNraHO21+f2hnH9+mvtcktGvE0Quttp886V+YXRtw8rv1xPsTsxo3SeK3IvKntfWI6wlcjXXumoj0sZPkmQsOfGTZICsCFW/pLJGOfE8pc0gUaKMafKmdaFrnrrI3CnJGEUJ+nG0ri7rBq2NEvT9dX3GaneARtT2XZvfv7acWW1SESibNFHJ74hU9ZMkGUtO/CQZIDnxk2SALEsbf5raZZbSDjyPtZW8HWXt4ii0MrL7+h7TaDyOPfbY4rFaF5gf01IuehvyCqO7zPyxkoutj5LffSTziGzwaexzux5idwnakueRHL6/2tLjEWnjJ0kylpz4STJAlqWqP+Z7B19Hap1XNy32XKvaehdVNB7W7WXl2L9//8h5fai2kSsrcsXZHW72WFT+yquvpbHy8kY55kq1CyLXYZSnzsrod/F12fHYRx0AGDUpoyjHPty/taSqnyTJWHLiJ8kAWRGqfolpSh31sdGiFN1lo9tgVK2rTeYBoyprSWWH2Cthr7OkhnqZfVKKUkkqr66WTB+oTy5hZfQRc6XvTbPxqVRWzXsyalVx7y0q3c9pNth03UxVIlX9JEnGkhM/SQZITvwkGSBTFc1cDlh7KCoL5SlFZk1jU9n+bF9RG9PIaG1Va6v779g2fXRXqQRYtDsvIkpQYW38KKmotf+9HNEuvloZrT3t3Xm2zdp75omSp5bKqk/jPi25NGfpAsxf/CQZIDW189aKyJdE5DoRuUFE3tp+fpKIXCMiN4vIR0Vk9ezFTZKkDya689qimUeq6oG2au5fAecCrwc+oaofEZH3Atep6nsmtNXJV2FdKHYjjseqSdZ1A6MbKKIcbbUbOSJqc/V5lc+68GwbXka7AaQ2X35tJKDHqvPeVWbH1EcvWjeadXlFrkmvppfU42mi7krX1tVt5l3IJTV9GndelxJxEb2487Rh4eoOa/8pcCbwsfbz7cDZVVIlSbLkVNn4InJoWyl3N3AV8G1gn6ourD7sBE4sfHebiOwQkR19CJwkyeKpmviq+pCqngJsBk4Hnj7utMJ3L1bV01T1tO5iJknSJ1O581R1n4h8AXgusF5EVrW/+puBO/sSKqrRZm3EaCdW5DaL3Dpdki76ME7r1vHtR3XYbKnpyOVobb/IBVa70zC6zmhMo8SQ9tjq1Y+s+0b57L1rsnQvonUNb4PbdZM+EllGdfVqd1v682yb9pmeZTh9zar+Y0Vkffv6cOAs4EbgauDl7WlbgctnJWSSJP1S84u/CdguIofS/KG4VFWvFJFvAB8RkbcBXwUumaGcSZL0yIrYnVfawWXdSTDq5urqCoki1UoJJXw0lz3m1WOr1tXmVIt2vvkxsK5Ke8zmx4O43Fgfz0RJtY2uxavwli4JO6BcVr1reapavMkRJenoWg67RO7OS5JkLDnxk2SArAhV37Vx8PU0amNpw0NUudSr+qUV+Wil2stoI9x8MoiSyjfN5pVSPr7alftpiHL/1aYi75KDcJpndtbtWyKPUN/l1yJS1U+SZCw58ZNkgOTET5IBsqJt/GlsX3tulJzRJqGMSmNF4xbZerZ930ZpfaF215rvL3LZWbqUkob6vPoRXWoL9FF2O6ol0MeOzWlKfPc9B9PGT5JkLDnxk2SArIice30kKii14V1qfUR0lXLzTcLmaYtcZdYk8JF7li7uNiiPo//cfi9yF1qiJBqeUmKSaINNtOnKl96yWHPKX6d9dnwNAvsMWjddJIdvY1559iz5i58kAyQnfpIMkJz4STJAVoSNX8qJ710yUX28UhtRuG1PdcxG3ke2Xims09t90bqBd/2V6Jpj3lJKIAFll+k0bj8rox2rrq640hoKxKW8LVGCVNt+5Ar2tRbttUV1Bvskf/GTZIDkxE+SAbLiIvdKJYs8kcunto0oYq6PnV5+Z2BJlq754brkD4Ru5k4UQWhNmEj2WrMoUudrx9u79iL1Psr9XzIvI5PAHyuVPe/q2svIvSRJxpITP0kGyIpT9V17xffRdZVW+P37aZJoWLquYpe8DbXXErUXfa+P6q1d+6qNIOw7UYYvxRaZPpZos1NkWm3cuPHga5//sO9cgKnqJ0kylpz4STJAcuInyQBZETZ+l4SJtdFd3p6rTaJRahtG7f9pbLaSnRxdS+1aQ5TY049pqY5BFDEYJZeM+qpdG6ilNlGGd9/V9tW1VHq0rpSJOJIkmQvVE78tlf1VEbmyfX+SiFwjIjeLyEdFZPWkNpIkWR5Ms0nnXJpimQs63DuAd6vqR0TkvcA5wHv6EKprmSUb+ebVqdrND1ESjWhzj6XWHPHqppXfRpl5eSOV26r31jTx8tq+fEKTUjSalzeqpGvvoe2rtkyWlytyJdZGKNox8OetX7++eMxSG0XZ1QU7L6p+8UVkM/CLwPva9wKcCXysPWU7cPYsBEySpH9qVf2LgDcCCz8FxwL7VHXhz/BO4MRxXxSRbSKyQ0R2LErSJEl6Y+LEF5FfAnar6rX24zGnjtVfVPViVT1NVU/rKGOSJD1TY+OfAfyyiLwUWEtj418ErBeRVe2v/mbgzr6EimzayD6yiRAid02te7A2QWWU/MLvwIvWK6JEEZYov78dq6678yz22mpLUMOoTW5de5Ft3TXJRbQuE4ULl9pYs2bNyDF7L2pLm8/aZbdYJv7iq+qbVHWzqm4BXgV8XlV/FbgaeHl72lbg8plJmSRJryzGj38+8HoRuYXG5r+kH5GSJJk1KyJyrwu15a+nSbZhVf9I3a7dpbV3796RY9Z91Xe+tWgXoldfS6ZQ1EbUn3XT+QjC2h1/tSp2LX1FDHYpSz7LXHqQkXtJkhTIiZ8kA+RRpepHKrZdPbbqpV+5t5FvvjyVVc1rVTe/klxbcbf0nUmU1PQoGtJfp00UYcfNmyb2e/5YbcRcF/Omjw0w3hSs3dAUeXrs97xp0repEpGqfpIkY8mJnyQDJCd+kgyQR5WNb1m9enSXcCkRh0+UUeuiqnXn+Vz5Ubnnki0cJaGMbPfIfra2am2E4jRuqJIrzvdVa/tG6zd2rcEnsvQ7D0ttWPvcPwP22fHt2XtfKhvmv1frwuw6N9PGT5JkLDnxk2SArIhqubVE6qDNI2fVbx+5Z1X/2rJQXnUrlUSCOBrN9mfPizb6eFdcKVlI7bVAeVNUlC8vMhdq3X5dE6SUxs3LZe9t1808UfSijQ6NyqP5Z8JSa+ItlvzFT5IBkhM/SQZITvwkGSCPWndeZItFbpcoKaf9Xslt5r8XhYZ6G65kL/rz7NqAT1DZJad/bQ0CT2SPlhJ2zjpBReTGrU246olcbKV6fL4Gg10Hqq2fmO68JEl6JSd+kgyQFafq96EKlVR26KcsdJSQwaqAvn0rl803568zcgdZ1d/urNuzZ8/IeZEbrUvCEX/Mqvpd71nt96JISXst1nXrx74Pc8SaGb4Mt21jmuQvXUhVP0mSseTET5IBsiJU/dIq/DT520pRfdHqa20prNrzPFEJLXvMR5lFMpauM0p/3bVkmTVH/HVaT8f+/fuL7VnPQ60HJDJNIiITzJpFfqOP7TuKtow8PfZ7frxL49M1Ui9V/SRJxpITP0kGSE78JBkgy9LGr92N1jUvfWR3R5RsxCg5o6d2PSDa+VaqEQCjrr6SS80fq5XDyx65RWvLTlub1rsc7fpLtGsyKo9ux8PeF39e12ux98aOqXe52rGLkoBE61S1pI2fJMlYqvbji8htwD3AQ8CDqnqaiGwEPgpsAW4DXqmqe0ttJEmyfKhS9duJf5qqft989k5gj6peKCIXABtU9fwJ7VSp+l6VK5Vg8uqqvZYuudwgVhttxFXknrHy1qqenlpzJIoys3L5zUiR26jkoorqDNSaOzYhipcrSjzRxTSBsmkV3TPviiuVA/PndnEFTzq3C7NW9V8GbG9fbwfOXkRbSZLMkdqJr8BnRORaEdnWfna8qu4CaP8/btwXRWSbiOwQkR2LFzdJkj6ozbl3hqreKSLHAVeJyDdrO1DVi4GLYb778ZMkKVM18VX1zvb/3SJyGXA6cJeIbFLVXSKyCdjdl1DePrS2dbQTq/Qd36a1sXxfa9euHXsejNpi1o1jQy4hdit22WU2zTpBya73Nnjk6rNt2mvxsvtQ4hL2e9F4eLdll1BWvxZQCvf26xWWyPUZrSvNohbirJio6ovIkSKybuE18CLgeuAKYGt72lbg8lkJmSRJv9T84h8PXNb+FVwFfEhV/0JEvgxcKiLnALcDr5idmEmS9MmyjNybor3ie39dpai7PnKXR3J41bDLrr7I5Ihy3Ze+A/WRaqXS4FBfysuaGdNENZbGKorstLJ7+aMIuWg8rLkTqfp9kDn3kiSZGTnxk2SA5MRPkgHyqKqdV+uiKtWG60pUN27jxo0jxyKbuRTyWesC9H1HLrCozS5j5du334tcsFFy01J/tWsBENvnFtu3dwX3YXdH1GZNytp5SZIsipz4STJAVoQ7r1Tu2e84i5J0lFxKXa8/2j1nZfS7uayMPvFEbSKRWrlqr82bKrU7G0v9empdcVEpMouXt1TGylO7M7IPGT2lWgX+fbrzkiSZGTnxk2SALMtVfa9qlTbjRPnm/bGSOhglkIii3Wo3CHnVMMrVV1Kxp0ncUKsellaSoX5DTLRab2UuJTCB0fsUmRhd+oJuZl10XmQWlZKgwGgNAk9U2mtW5C9+kgyQnPhJMkBy4ifJAFmW7rwoEWIp1zrU23Al9yDEyR9rZBonV4naiDzvEuxSynua5I9d6w6U6OJug1H7vLY+Xtf1j9oah9HzUhslWNtG1s5LkqRXcuInyQBZlqq+p6TmdYkwm0Tf5a+nKZdU6tvnoq8tY11SIWFUjYzcp14tLeHVUiuXjaL0ZlGkztpxtO4wf1+8C89Sa7pF9QPsMT8e9hm07fuo0kjePiJJLanqJ0kylpz4STJAcuInyQBZETZ+KTf6NIkKZp1MoUTXxJDWzty3b9/IeTa5h7clu5Rt9u5Ca4fXupRq3Whdk37aHPvrXeiUAAAFUElEQVRRqLMfj9oS1KVxgzjxaSmxalSGe5o6DF1IGz9JkrHkxE+SAbIsVf0ossmqf1GOtqiN2jzvs6BLVFxtXj0ouzgj12ftbkjfV8k08cdqS3JHEYqWPvLeTzOmtfUDat2F3lTpe0deqvpJkoylauKLyHoR+ZiIfFNEbhSR54nIRhG5SkRubv8vVyFMkmRZUaXqi8h24C9V9X0isho4AngzsEdVLxSRC4ANqnr+hHYWnXPPMo2ZslSr+rNg1tdSWsWeZgNMFxlrE45Mk5iklq5j2iXSc9bPX42qP3Hii8jRwHXAE9WcLCI3AS8wZbK/oKpPndBWTvweyImfEz+iLxv/icD3gD8Rka+KyPvactnHq+qutqNdwHHjviwi20Rkh4jsmEL2JElmSM3EXwWcCrxHVZ8F3AtcUNuBql6sqqep6mkdZUySpGdqkm3uBHaq6jXt+4/RTPy7RGSTUfV3z0rInnYs9SDJ8mAOquKi++ryvT6SYXZl1te53J6/ib/4qvpd4A4RWbDffx74BnAFsLX9bCtw+UwkTJKkd2pX9U8B3gesBm4FXkPzR+NS4PHA7cArVHVPsRG6L+4lSVJPL6v6fZITP0lmT0buJUkylpz4STJAcuInyQDJiZ8kAyQnfpIMkJz4STJA5l0m+/vAd4DHtK+XkuUgA6QcnpRjlGnleELNSXP14x/sVGTHUsfuLwcZUo6UY6nkSFU/SQZITvwkGSBLNfEvXqJ+LctBBkg5PCnHKDORY0ls/CRJlpZU9ZNkgOTET5IBMteJLyIvFpGbROSWNjPvvPp9v4jsFpHrzWdzTw8uIo8TkavbFOU3iMi5SyGLiKwVkS+JyHWtHG9tPz9JRK5p5fhom1F55ojIoW0+xyuXSg4RuU1Evi4iX1vID7lEz8hcUtnPbeKLyKHAHwAvAZ4BvFpEnjGn7v8UeLH77ALgc6r6ZOBzTJFHcBE8CPymqj4deC7w2nYM5i3LPwJnqupPAacALxaR5wLvAN7dyrEXOGfGcixwLnCjeb9Ucvycqp5i/OZL8Yz8T+AvVPVpwE/RjEv/cqjqXP4BzwM+bd6/CXjTHPvfAlxv3t8EbGpfbwJumpcsRobLgRcupSw0NRK+Avw0TYTYqnH3a4b9b24f5jOBKwFZIjluAx7jPpvrfQGOBv6edtF9lnLMU9U/EbjDvN/ZfrZUVKUHnxUisgV4FnDNUsjSqtdfo0mSehXwbWCfqi4UcpvX/bkIeCOwUPTu2CWSQ4HPiMi1IrKt/Wze92VRqeynYZ4Tf1w6oEH6EkXkKODjwHmqevek82eBqj6kqqfQ/OKeDjx93GmzlEFEfgnYrarX2o/nLUfLGap6Ko0p+loR+Zk59OlZVCr7aZjnxN8JPM683wzcOcf+PXe1acGZdXpwi4gcRjPpP6iqn1hKWQBUdR/wBZo1h/UisrBxax735wzgl0XkNuAjNOr+RUsgB6p6Z/v/buAymj+G874v41LZnzoLOeY58b8MPLldsV0NvIomRfdSMff04NLUUboEuFFV37VUsojIY0Vkffv6cOAsmkWkq4GXz0sOVX2Tqm5W1S00z8PnVfVX5y2HiBwpIusWXgMvAq5nzvdF55nKftaLJm6R4qXAt2jsyf84x34/DOwCHqD5q3oOjS35OeDm9v+Nc5Dj+TRq698BX2v/vXTesgA/CXy1leN64D+3nz8R+BJwC/BnwJo53qMXAFcuhRxtf9e1/25YeDaX6Bk5BdjR3ptPAhtmIUeG7CbJAMnIvSQZIDnxk2SA5MRPkgGSEz9JBkhO/CQZIDnxk2SA5MRPkgHy/wGEqAQniyu29wAAAABJRU5ErkJggg==\n", + "text/plain": [ + "
" + ] + }, + "metadata": {}, + "output_type": "display_data" } ], "source": [ diff --git a/README.md b/README.md index 4564651..39f1731 100644 --- a/README.md +++ b/README.md @@ -1,2 +1,55 @@ # CoordConv +![](https://img.shields.io/badge/pytorch-0.4.0-blue.svg) ![](https://img.shields.io/badge/python-3.6.5-brightgreen.svg) Pytorch implementation of CoordConv for N-D ConvLayers, and the experiments. + +Reference from the paper "An intriguing failing of convolutional neural networks and the CoordConv solution." + +Extends the CoordinateChannel concatenation from 2D to 1D and 3D tensors. + +# Requirements +- pytorch 0.4.0 +- torchvision 0.2.1 +- torchsummary 1.3 +- sklearn 0.19.1 + +# Usage +```python +from coordconv import CoordConv1d, CoordConv2d, CoordConv3d + +class Net(nn.Module): + def __init__(self): + super(Net, self).__init__() + self.coordconv = CoordConv2d(2, 32, 1, with_r=True) + self.conv1 = nn.Conv2d(32, 64, 1) + self.conv2 = nn.Conv2d(64, 64, 1) + self.conv3 = nn.Conv2d(64, 1, 1) + self.conv4 = nn.Conv2d( 1, 1, 1) + + def forward(self, x): + x = self.coordconv(x) + x = F.relu(self.conv1(x)) + x = F.relu(self.conv2(x)) + x = F.relu(self.conv3(x)) + x = self.conv4(x) + x = x.view(-1, 64*64) + return x + +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") +net = Net().to(device) +``` + +# Experiments +Implement experiments from origin paper. + +## Coordinate Classification +Use `experiments/generate_data.py` to generate `Uniform` and `Quadrant` datasets for Coordinate Classification task. + +Use `experiments/train_and_test.py` to train and test neural network model. + +### Images + +|Train|Test|Predictions| +|:---:|:---:|:---:| +|![](https://i.loli.net/2018/07/16/5b4c7db11abf9.png)|![](https://i.loli.net/2018/07/16/5b4c7dbd03169.png)|![](https://i.loli.net/2018/07/16/5b4c8d88a70a2.png)| + +