Created
August 25, 2017 03:23
-
-
Save justheuristic/642bf3240f685b0372d48ae43ae9418f to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| { | |
| "cells": [ | |
| { | |
| "cell_type": "code", | |
| "execution_count": 1, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "import numpy as np\n", | |
| "import theano,theano.tensor as T\n", | |
| "import lasagne\n", | |
| "from lasagne.layers import *" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### Motivation, simple\n", | |
| "\n", | |
| "Sometimes it's easier to point out what's wrong than to define what's right\n", | |
| "* How humans evaluate images, translations, etc\n", | |
| "* How \"slow\" art style transfer is better than \"fast\"\n", | |
| "* GANs vs likelihood-based models\n", | |
| "\n", | |
| "So let's define a network that contains a layer that does energy minimization to get rid of errors in the image.\n", | |
| "\n", | |
| "Possible applications:\n", | |
| "* MT - fixing \"stupid\" translations\n", | |
| "* GANs - converting discriminator into better generator\n", | |
| "* Segmentation / colorization: fixing implausible outputs" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 2, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "class ConstraintLayer(Layer):\n", | |
| " \"\"\"\n", | |
| " A lasagne layer that performs a minimization operation w.r.t. a learned objective.\n", | |
| " Trains end-to-end with gradient descent.\n", | |
| " It minimizes a mean of learned cost functions.\n", | |
| " If all cost functions have a lower bound, this is effectively equivalent to solving CSPs.\n", | |
| " \"\"\"\n", | |
| " def build_inner_net(self,input_shape,num_units,\n", | |
| " nonlinearity=lasagne.nonlinearities.elu):\n", | |
| " \n", | |
| " l_in = InputLayer(input_shape)\n", | |
| " l_hid = DenseLayer(l_in,num_units,nonlinearity=nonlinearity)\n", | |
| " l_out = ExpressionLayer(l_hid,lambda a: a.mean(-1),output_shape=input_shape[:1])\n", | |
| " return l_out\n", | |
| " \n", | |
| " \n", | |
| " def __init__(self,incoming,num_units=100,inner_net=None,\n", | |
| " n_steps=1,learning_rate=1,name=None,\n", | |
| " **kwargs):\n", | |
| " self.n_steps = n_steps\n", | |
| " self.learning_rate = learning_rate\n", | |
| " self.inner_net = inner_net or self.build_inner_net(incoming.output_shape,\n", | |
| " num_units,**kwargs)\n", | |
| " \n", | |
| " Layer.__init__(self,incoming,name=name)\n", | |
| " \n", | |
| " def get_output_for(self,input,**kwargs):\n", | |
| " \n", | |
| " def step(input):\n", | |
| " objective = get_output(self.inner_net,input)\n", | |
| " \n", | |
| " grads = T.grad(objective.mean(),input)\n", | |
| " return input - self.learning_rate*grads\n", | |
| " \n", | |
| " if self.n_steps == 1:\n", | |
| " final_input = step(input)\n", | |
| " else:\n", | |
| " history,updates = theano.scan(step,outputs_info=[input],\n", | |
| " n_steps=self.n_steps)\n", | |
| " \n", | |
| " final_input = history[-1]\n", | |
| " \n", | |
| " assert len(updates)==0,\"Default updates for objectives are not supported\"\n", | |
| " \n", | |
| " return final_input\n", | |
| " def get_params(self,**kwargs):\n", | |
| " return get_all_params(self.inner_net,**kwargs)" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### Toy problem: train a network to produce numbers in increments of 1\n", | |
| "\n", | |
| "This time we only build simple constraint layer that does not condition on anything. We want it to produce sequences of incrementing numbers from random noise.\n", | |
| "\n", | |
| "As a sanity check, i want to see that the network doesn't learn to always output same numbers (as DenseLayer would likely do in this case)." | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 3, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "l_in = InputLayer((None,5))\n", | |
| "l_out = ConstraintLayer(l_in,n_steps=50,learning_rate=10)\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 4, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "x = T.matrix()\n", | |
| "y_pred = get_output(l_out,x)\n", | |
| "\n", | |
| "loss = T.mean((y_pred[:,:-1] - y_pred[:,1:] +1)**2)\n", | |
| "\n", | |
| "updates = lasagne.updates.adadelta(loss,get_all_params(l_out))\n", | |
| "train_step = theano.function([x],loss,updates=updates,allow_input_downcast=True)\n", | |
| "predict_fun = theano.function([x],get_output(l_out,x),allow_input_downcast=True)" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 5, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "name": "stderr", | |
| "output_type": "stream", | |
| "text": [ | |
| "/usr/local/lib/python2.7/dist-packages/matplotlib/font_manager.py:273: UserWarning: Matplotlib is building the font cache using fc-list. This may take a moment.\n", | |
| " warnings.warn('Matplotlib is building the font cache using fc-list. This may take a moment.')\n", | |
| "100%|██████████| 10000/10000 [01:38<00:00, 101.23it/s]\n" | |
| ] | |
| }, | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "[<matplotlib.lines.Line2D at 0x7fdcf0960c90>]" | |
| ] | |
| }, | |
| "execution_count": 5, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| }, | |
| { | |
| "data": { | |
| "image/png": "iVBORw0KGgoAAAANSUhEUgAAAYAAAAEACAYAAAC6d6FnAAAABHNCSVQICAgIfAhkiAAAAAlwSFlz\nAAALEgAACxIB0t1+/AAAFAZJREFUeJzt3WuwXWV9x/HvP5xEETQECsEmmHBTBxRSdEIUZtgtgmHq\nwLSjBacdLvWFgxccZSzqOCa8q87UoqWIdNACKmCplYBSQoXNRbmHGAoJpCB3CUgIEi6ZXP59sdbh\nbE72yd4J+xLyfD8ze1h7rWev51nPXpzfetZlJzITSVJ5Jg27AZKk4TAAJKlQBoAkFcoAkKRCGQCS\nVCgDQJIK1TEAIuItEXF7RNwTEfdGxII2ZaZExGURsTIibo2Id/WnuZKkXukYAJm5DvjzzPwzYA5w\nXETMHVfsU8DqzDwQOAf4Vs9bKknqqa5OAWXmy/XkW4ARYPzTYycAF9XTVwBH96R1kqS+6SoAImJS\nRNwDPA1cl5l3jisyA3gcIDM3AmsiYveetlSS1FPdjgA21aeAZgKHR8RB44pEm/f+xoQkbcdGtqZw\nZv4xIprAfOD+lkWPA/sAT0XETsA7MvP58Z+PCENBkrZBZo4/0H7DurkL6E8iYmo9vTPwEWDFuGJX\nAafU058Arp9ofZnpK5MFCxYMvQ3by8u+sC/siy2/+qWbEcA7gYsiYhJVYFyemb+MiLOBOzPzauBC\n4JKIWAk8B5zUtxZLknqiYwBk5r3AYW3mL2iZXgf8TW+bJknqJ58EHpJGozHsJmw37Isx9sUY+6L/\nop/nlzarLCIHWZ8k7QgighzGRWBJ0o7JAJCkQhkAklQoA0CSCmUASFKhDABJKpQBIEmFMgAkqVAG\ngCQVygCQpEIZAJJUKANAkgplAEhSoQwASSqUASBJhTIAJKlQBoAkFcoAkKRCGQCSVCgDQJIKZQBI\nUqEMAEkqlAEgSYUyACSpUAaAJBXKAJCkQnUMgIiYGRHXR8T9EXFvRJzRpsxREbEmIpbUr6/3p7mS\npF4Z6aLMBuBLmbk0InYF7o6IxZm5Yly5mzLz+N43UZLUDx1HAJn5dGYurafXAsuBGW2KRo/bJknq\no626BhARs4E5wO1tFs+LiHsi4hcRcVAP2iZJ6qNuTgEBUJ/+uQL4Qj0SaHU3MCszX46I44CfA+/u\nXTMlSb3WVQBExAjVH/9LMvPK8ctbAyEzr4mI8yJi98xcPb7swoULX5tuNBo0Go1taLYk7biazSbN\nZrPv9URmdi4UcTHwh8z80gTLp2fmqnp6LvDTzJzdplx2U58kaUxEkJk9v87acQQQEUcAfwvcGxH3\nAAl8DZgFZGZeAHw8Ik4H1gOvACf2uqGSpN7qagTQs8ocAUjSVuvXCMAngSWpUAaAJBXKAJCkQhkA\nklQoA0CSCmUASFKhDABJKpQBIEmFMgAkqVAGgCQVygCQpEIZAJJUKANAkgplAEhSoQwASSqUASBJ\nhTIAJKlQBoAkFcoAkKRCGQCSVCgDQJIKZQBIUqEMAEkqlAEgSYUyACSpUAaAJBXKAJCkQhkAklQo\nA0CSCtUxACJiZkRcHxH3R8S9EXHGBOW+GxErI2JpRMzpfVMlSb000kWZDcCXMnNpROwK3B0RizNz\nxWiBiDgO2D8zD4yIw4HzgXn9abIkqRc6jgAy8+nMXFpPrwWWAzPGFTsBuLguczswNSKm97itkqQe\n2qprABExG5gD3D5u0Qzg8Zb3T7J5SEiStiPdnAICoD79cwXwhXok8LrFbT6S7dazcOHC16YbjQaN\nRqPbJkhSEZrNJs1ms+/1RGbbv9OvLxQxAlwNXJOZ32mz/Hzghsy8vH6/AjgqM1eNK5fd1CdJGhMR\nZGa7A+03pNtTQD8A7m/3x7+2CDgZICLmAWvG//GXJG1fOo4AIuII4CbgXqrTOgl8DZgFZGZeUJc7\nF5gPvASclplL2qzLEYAkbaV+jQC6OgXUs8oMAEnaasM+BSRJ2sEYAJJUKANAkgplAEhSoQwASSqU\nASBJhTIAJKlQBoAkFcoAkKRCGQCSVCgDQJIKZQBIUqEMAEkqlAEgSYUyACSpUAaAJBXKAJCkQhkA\nklQoA0CSCmUASFKhDABJKpQBIEmFMgAkqVAGgCQVygCQpEINPAAyB12jJKmdgQfAI48MukZJUjue\nApKkQnUMgIi4MCJWRcSyCZYfFRFrImJJ/fp675spSeq1kS7K/BD4F+DiLZS5KTOP702TJEmD0HEE\nkJm3AM93KBa9aY4kaVB6dQ1gXkTcExG/iIiDerROSVIfdXMKqJO7gVmZ+XJEHAf8HHj3RIXPOWch\n06ZV041Gg0aj0YMmSNKOo9ls0mw2+15PZBc35kfELOCqzDyki7K/Az6QmavbLMuHH0723Xeb2ipJ\nRYoIMrPnp9q7PQUUTHCePyKmt0zPpQqVzf74S5K2Lx1PAUXET4AGsEdEPAYsAKYAmZkXAB+PiNOB\n9cArwIlbWp9PAkvS9qGrU0A9qywiH3oo2W+/gVUpSW96wz4F1DPhDaOStF3wx+AkqVADD4ANGwZd\noySpnYEHwI9/POgaJUntDDwAXn110DVKktrx56AlqVAGgCQVygCQpEJ5G6gkFcoRgCQVyieBJalQ\nngKSpEIZAJJUKK8BSFKhHAFIUqEMAEkqlKeAJKlQjgAkqVAGgCQVygCQpEIZAJJUKANAkgplAEhS\nobwNVJIK5QhAkgo18ACYPHnQNUqS2hl4AHzsY4OuUZLUjv8gjCQVqmMARMSFEbEqIpZtocx3I2Jl\nRCyNiDm9baIkqR+6GQH8EPjoRAsj4jhg/8w8EPg0cP6WVuZFYEnaPnQMgMy8BXh+C0VOAC6uy94O\nTI2I6b1pniSpX3pxDWAG8HjL+yfreW05ApCk7cNID9bR7rLuhH/mL754IbfcUk03Gg0ajUYPmiBJ\nO45ms0mz2ex7PZFdHJJHxCzgqsw8pM2y84EbMvPy+v0K4KjMXNWmbC5enBxzzBtvuCSVIiLIzJ7f\nQ9ntKaCg/ZE+wCLgZICImAesaffHX5K0fek4AoiInwANYA9gFbAAmAJkZl5QlzkXmA+8BJyWmUsm\nWFdCeh1AkrZCv0YAXZ0C6lllBoAkbbVhnwKSJO1gDABJKpQBIEmFMgAkqVAGgCQVygCQpEINJQCe\ne24YtUqSWg3lOQDwR+EkqVs+ByBJ6ikDQJIKZQBIUqEMAEkqlAEgSYUyACSpUAaAJBXKAJCkQhkA\nklQoA0CSCmUASFKhDABJKpQBIEmFMgAkqVAGgCQVygCQpEIZAJJUKANAkgplAEhSoQwASSpUVwEQ\nEfMjYkVEPBgRZ7VZfkpEPBMRS+rX3/e+qZKkXhrpVCAiJgHnAkcDTwF3RsSVmbliXNHLMvOMPrRR\nktQH3YwA5gIrM/PRzFwPXAac0KZc9LRlkqS+6iYAZgCPt7x/op433l9HxNKI+GlEzOxJ6yRJfdNN\nALQ7ss9x7xcBszNzDvAr4KJOK122rIuaJUl90/EaANUR/7ta3s+kuhbwmsx8vuXtvwHfnHh1CwE4\n9FC44YYGjUajq4ZKUimazSbNZrPv9UTm+IP5cQUidgIeoLoI/HvgDuCTmbm8pczemfl0Pf1XwJcz\n88Nt1pWtg4cOVUuSgIggM3t+nbXjCCAzN0bE54DFVKeMLszM5RFxNnBnZl4NnBERxwPrgdXAqb1u\nqCSptzqOAHpamSMASdpq/RoB+CSwJBVqqAHgCECShmeoAXDFFcOsXZLKNtQA+M1vhlm7JJVtqAGw\nYcMwa5eksg01ADZtGmbtklQ2LwJLUqEcAUhSoYb6IBg4CpCkTnwQTJLUUwaAJBVq4AFw8MGDrlGS\n1M7AA+DSSwddoySpnYEHwMaNg65RktTOwAMgYsvvJUmD4UVgSSqUASBJhTIAJKlQQ78GIEkajoEH\nwPvet/m8z3xm0K2QJA38t4Ays+0owN8EkqT2dqjfAjr55GHUKklqNZQA+PznN583Z87g2yFJJRtK\nAEyZsvm83/528O2QpJINJQD237/9/OXL4eab4aSTBtseSSrRUAJgl13az7/xxurH4i6/fLDtkaQS\nDe1BsHZ3/Zx++uY/FrduHcyYMZg2SVJJhnIb6KiNG2FkpH3Z0WLPPAPTp3ubqKRyDfU20IiYHxEr\nIuLBiDirzfIpEXFZRKyMiFsj4l3drHennSZedu21sHq1f/glqV86BkBETALOBT4KHAx8MiLeO67Y\np4DVmXkgcA7wrW4bcMwx7efPnw977AELFlTvd7QgaDabw27CdsO+GGNfjLEv+q+bEcBcYGVmPpqZ\n64HLgBPGlTkBuKievgI4utsGLF4MmzZNvPz736/+u27d2LwNG6pXt159FX73u+7LD4I79xj7Yox9\nMca+6L9uAmAG8HjL+yfqeW3LZOZGYE1E7N5tIyKqI/zZsycus/POVbkImDy5eo1Of/azsOuucOyx\n8O1vwyOPwEsvwZIl8NBDcOaZsN9+1XpGw+bll8dGFaMXnltHGZnV/GGOPNau3Twc16+v2rSjjYgk\nDd4El2Bfp92Fh/F/fsaXiTZlOho9Sr/vPvjZz+Ab3+j8mQ0b4LzzqunrrqteZ57Zvuywf4l0772r\noHrHO+CJJ+Dss6v5++0HDz9ctW/KFHjnO2HmTLjllmr54YdXF8t//evN1zl5chUKhxwCy5ZV61+7\ndmz5AQdU9U2aBCtXVhfUH3oIjjyyugazfn317AXA0UdXobdpUzVquuOOav573gMPPAAHH1xdlH/2\nWfjgB6vPv+1t1XcwZUoVui+/DHvuCY8+Wm3vLrtU61u3rtq+V16BVatgn32qzwI8+CD86EdVPzz/\nPNx1V9WWm2+GefOqdUZU65o8uWrjXXfB3LljYbh2LTz5JBx00Ni2b9xYtW20Xzdtql4jI1V/bNhQ\nbUNm9d/RA4Fnn60+s/POVdmNG+Htbx8L3dH/rllT9dGHPlRt3557VvMXL4Z994UDD9z8+9q0CV58\nEW67rXrfaFSfffLJsRsiRpetWVN9D6++Wt0JNzKy+T48uv233FLVN3rH3Oj8TZs2n548udr+F1+s\n9pfRdS5bVj2js+uuVdnR+ddeW+1fe+0FN91U7VO77ALTplX7wzPPwPvfv/m2TvT/2+hBzOTJY+2a\nNGls+le/qtY9bdpYX0yaVC278UY46qixdbfW8dhjsPvuVdtGt3l0vaMHj6P7QMTYd996MPXUU9Wp\n57e+tXr/wgvV+kZGNj9AbLe9EVWdy5dX+8Puu48ty4Rbb63Wd+ih1fx168a+99YbYp59tvo7eOSR\n1fr6peNdQBExD1iYmfPr918BMjO/2VLmmrrM7RGxE/D7zNyrzbo8bpWkbdCPu4C6GQHcCRwQEbOA\n3wMnAZ8cV+Yq4BTgduATwPXtVtSPDZAkbZuOAZCZGyPic8BiqmsGF2bm8og4G7gzM68GLgQuiYiV\nwHNUISFJ2o4N9EEwSdL2Y2A/BdHpYbI3u4iYGRHXR8T9EXFvRJxRz58WEYsj4oGIuDYiprZ85rv1\nw3NLI2JOy/xT6n56ICLetP96QkRMioglEbGofj87Im6rt+vSiBip50/4IGFEfLWevzwijh3WtrwR\nETE1Iv6j3ob7IuLwUveLiPhiRPxvRCyLiB/X330R+0VEXBgRqyJiWcu8nu0HEXFY3a8PRsQ5XTUq\nM/v+ogqa/wNmAZOBpcB7B1H3oF7A3sCcenpX4AHgvcA3gX+o558F/GM9fRzwi3r6cOC2enoa8BAw\nFdhtdHrY27eNffJF4EfAovr95cAn6unvAZ+up08HzqunTwQuq6cPAu6hOlU5u96HYtjbtQ398O/A\nafX0SP3dFrdfAH8KPAxMadkfTillvwCOBOYAy1rm9Ww/oLoGO7ee/iXw0Y5tGtCGzwOuaXn/FeCs\nYX8hfd7mnwMfAVYA0+t5ewPL6+nzgRNbyi8HplNdP/ley/zvtZZ7s7yAmcB1QIOxAHgWmDR+nwD+\nGzi8nt4JeKbdfgJcM1ruzfIC3g481GZ+cftFHQCP1n/ERoBFwDHAM6XsF1QHwa0B0JP9oP7s/S3z\nX1duotegTgF18zDZDiMiZlMl/W1UX+4qgMx8Ghi9PXaiPhk//0nenH31z8CXqZ8HiYg9gOczc/TR\nttZ9YPyDhC/UDxLuCH2xH/CHiPhhfTrsgoh4GwXuF5n5FPBPwGNU7X8BWAKsKXC/GLVXj/aDGXWZ\n8eW3aFAB0M3DZDuEiNiV6ucwvpCZa5l4Oyd6eO5N31cR8ZfAqsxcytj2BJtvW7YsG2+H6AuqI93D\ngH/NzMOAl6iOYEvcL3aj+tmYWVSjgV2oTnWMV8J+0cnW7gfb1CeDCoAngNZfCJ0JPDWgugemvnh1\nBXBJZl5Zz14VEdPr5XtTDXeh6pN9Wj4+2ic7Ql8dARwfEQ8DlwJ/QfUjgVPrHxeE12/Xa31RP0g4\nNTOfZ+I+ejN5Ang8M++q3/8nVSCUuF98BHg4M1fXR/T/BXwY2K3A/WJUr/aDbeqTQQXAaw+TRcQU\nqvNTiwZU9yD9gOo83Hda5i0CTq2nTwWubJl/Mrz2tPWaeih4LXBMfefINKpzpNf2v+m9k5lfy8x3\nZeZ+VN/19Zn5d8ANVA8KQnXxr7UvTqmnWx8kXAScVN8Nsi9wAHDHILahV+rv9PGIeHc962jgPgrc\nL6hO/cyLiLdGRDDWFyXtF+NHwj3ZD+rTR3+MiLl1357csq6JDfDix3yqO2NWAl8Z9sWYPmzfEcBG\nqjuc7qE6tzkf2B34n3rbrwN2a/nMuVR3MPwWOKxl/ql1Pz0InDzsbXuD/XIUYxeB96W6U+FBqjs/\nJtfz3wL8tN7m24DZLZ//at1Hy4Fjh70929gHh1IdBC0FfkZ1B0eR+wWwoP4ul1H9gvDkUvYL4CdU\nR+XrqMLwNKoL4j3ZD4APAPfWy77TTZt8EEySCjW0fxNYkjRcBoAkFcoAkKRCGQCSVCgDQJIKZQBI\nUqEMAEkqlAEgSYX6fzIOTNtpxRwgAAAAAElFTkSuQmCC\n", | |
| "text/plain": [ | |
| "<matplotlib.figure.Figure at 0x7fdcf0fad450>" | |
| ] | |
| }, | |
| "metadata": {}, | |
| "output_type": "display_data" | |
| } | |
| ], | |
| "source": [ | |
| "from matplotlib import pyplot as plt\n", | |
| "from tqdm import trange\n", | |
| "%matplotlib inline\n", | |
| "\n", | |
| "history = [train_step(np.random.randn(10,5)) for _ in trange(10000)]\n", | |
| "plt.plot(history)" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### Sample output" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 6, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "array([[-1.95008159, -0.94695568, 0.04897024, 1.04185927, 2.04386783],\n", | |
| " [-1.9805001 , -0.96903646, 0.04872648, 1.06052589, 2.07302523],\n", | |
| " [-2.33256364, -1.34025693, -0.3424964 , 0.64436758, 1.64589214],\n", | |
| " [-2.34208632, -1.33263719, -0.33124033, 0.67538762, 1.68576896],\n", | |
| " [-2.61089253, -1.59195876, -0.56573665, 0.45609391, 1.47855389],\n", | |
| " [-1.9408406 , -0.94076324, 0.05987605, 1.05671751, 2.0595293 ],\n", | |
| " [-1.8643502 , -0.86960185, 0.14541446, 1.13533056, 2.15676403],\n", | |
| " [-1.42135119, -0.41845196, 0.58740461, 1.58860159, 2.592623 ],\n", | |
| " [-1.55481076, -0.53160506, 0.49015826, 1.50303304, 2.53067541],\n", | |
| " [-2.17136192, -1.17628527, -0.19068234, 0.78746402, 1.76838768]], dtype=float32)" | |
| ] | |
| }, | |
| "execution_count": 6, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| } | |
| ], | |
| "source": [ | |
| "_x = np.random.randn(10,5)\n", | |
| "predict_fun(_x)" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 7, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "array([[-1.00312591, -0.9959259 , -0.99288905, -1.00200856],\n", | |
| " [-1.01146364, -1.0177629 , -1.01179945, -1.01249933],\n", | |
| " [-0.99230671, -0.99776053, -0.98686397, -1.00152457],\n", | |
| " [-1.00944912, -1.00139689, -1.00662792, -1.01038134],\n", | |
| " [-1.01893377, -1.02622211, -1.02183056, -1.02245998],\n", | |
| " [-1.00007737, -1.00063932, -0.99684143, -1.00281179],\n", | |
| " [-0.99474835, -1.01501632, -0.98991609, -1.02143347],\n", | |
| " [-1.00289917, -1.00585651, -1.00119698, -1.00402141],\n", | |
| " [-1.02320576, -1.02176332, -1.01287484, -1.02764237],\n", | |
| " [-0.99507666, -0.98560292, -0.97814637, -0.98092365]], dtype=float32)" | |
| ] | |
| }, | |
| "execution_count": 7, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| } | |
| ], | |
| "source": [ | |
| "predict_fun(_x)[:,:-1] - predict_fun(_x)[:,1:]" | |
| ] | |
| } | |
| ], | |
| "metadata": { | |
| "kernelspec": { | |
| "display_name": "Python 2", | |
| "language": "python", | |
| "name": "python2" | |
| }, | |
| "language_info": { | |
| "codemirror_mode": { | |
| "name": "ipython", | |
| "version": 2 | |
| }, | |
| "file_extension": ".py", | |
| "mimetype": "text/x-python", | |
| "name": "python", | |
| "nbconvert_exporter": "python", | |
| "pygments_lexer": "ipython2", | |
| "version": "2.7.6" | |
| } | |
| }, | |
| "nbformat": 4, | |
| "nbformat_minor": 0 | |
| } |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment