Skip to content

Instantly share code, notes, and snippets.

@justheuristic
Created October 25, 2016 23:26
Show Gist options
  • Select an option

  • Save justheuristic/9b390548f4138ac8d1c14fde378485d4 to your computer and use it in GitHub Desktop.

Select an option

Save justheuristic/9b390548f4138ac8d1c14fde378485d4 to your computer and use it in GitHub Desktop.
Display the source blob
Display the rendered blob
Raw
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"collapsed": true
},
"outputs": [],
"source": [
"import theano\n",
"import numpy as np\n",
"from theano import tensor as T"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {
"collapsed": false
},
"outputs": [],
"source": [
"actual_y = T.vector()\n",
"predicted_y = T.vector()\n",
"query_id = T.vector()\n",
"\n",
"qi_equals_qj = T.eq(query_id[:,None], query_id[None,:])\n",
"\n",
"yi_lessthan_yj = actual_y[:,None] < actual_y[None,:]\n",
"\n",
"loss_i_j = T.exp(predicted_y[:,None] - predicted_y[None,:]) #or any other loss you think of\n",
"\n",
"loss = T.sum(qi_equals_qj * yi_lessthan_yj * loss_i_j)\n",
"\n",
"grad = T.grad(loss,predicted_y)\n",
"\n",
"hess = T.diagonal(T.hessian(loss,predicted_y))\n",
"\n",
"get_loss_and_derivatives = theano.function([predicted_y,actual_y,query_id],[loss,grad,hess])"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Sanity check:"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {
"collapsed": false
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"L = 60.995\n",
"L = 14.708\n",
"L = 5.410\n",
"L = 2.407\n",
"L = 1.147\n",
"L = 0.564\n",
"L = 0.282\n",
"L = 0.142\n",
"L = 0.072\n",
"L = 0.037\n",
"[ 1.20087573e-04 6.51633115e-02 0.00000000e+00 3.95972825e-01\n",
" 3.95969812e-01 3.95976326e-01 7.07454456e-01 7.07450525e-01\n",
" 1.00000000e+00 4.00683058e-02 4.78740161e-01 9.34306279e-01\n",
" 9.75975944e-01]\n"
]
}
],
"source": [
"#in action\n",
"\n",
"actual_y = np.array([0,0,0,1,1,1,2,2,3,0,1,2,2],'float32')\n",
"query_id = np.array([1,1,1,1,1,1,1,1,1,2,2,2,2],'float32')\n",
"#start with random guesses\n",
"predicted_y = np.random.randn(len(actual_y))\n",
"for i in range(100):\n",
" \n",
" L,dL,ddL = get_loss_and_derivatives(predicted_y,actual_y,query_id)\n",
" \n",
" if i%10==0:\n",
" print 'L = %.3f'%L\n",
" \n",
" predicted_y = predicted_y - 0.1 * dL / ddL\n",
" \n",
"\n",
"#convert all values into [0,1] for human readability\n",
"humanize = lambda a: (a - a.min()) / (a.max() - a.min())\n",
"print humanize(predicted_y)"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {
"collapsed": false
},
"outputs": [
{
"data": {
"text/plain": [
"<matplotlib.collections.PathCollection at 0x7fc98710dbd0>"
]
},
"execution_count": 4,
"metadata": {},
"output_type": "execute_result"
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/jheuristic/anaconda2/lib/python2.7/site-packages/matplotlib/collections.py:590: FutureWarning: elementwise comparison failed; returning scalar instead, but in the future will perform elementwise comparison\n",
" if self._edgecolors == str('face'):\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAX4AAAEACAYAAAC08h1NAAAABHNCSVQICAgIfAhkiAAAAAlwSFlz\nAAALEgAACxIB0t1+/AAAFX9JREFUeJzt3XuQXOV55/HvI40uSAhkRWsuEiBAskE22EBJyCvWNMZe\nD6TKcjmxWWzFcWw21NbiOFveCpCtDVOb1CqQqg2LqXALJq5CixxAcTDLAoKl1wQLkDA3B0mRLDAS\nYCxhbGMLkMQ8+8c0cjNoRjN9Rt0zer+fqinO5enzPjP0+c3R2316IjORJJVjXKcbkCS1l8EvSYUx\n+CWpMAa/JBXG4Jekwhj8klSYysEfEd+MiJcj4ukB9n8hIp6MiKci4qGIOLnqmJKk1o3EFf9NQPcg\n+zcDH83Mk4E/B64fgTElSS2qHPyZ+SDw6iD7V2fmLxqrjwCzq44pSWpdu+f4vwLc1eYxJUlNuto1\nUEScBXwZWNyuMSVJ79aW4G+8oHsD0J2Z75oWigg/MEiSWpCZMdzH7Pepnog4GlgJLM3MTQPVZeaY\n/brssss63oP9d74P+x97X2O598zWr5crX/FHxC3AmcDMiNgCXAZMaIT5dcCfAe8BrokIgF2ZubDq\nuJKk1lQO/sw8fx/7LwAuqDqOJGlkeOfuCKjVap1uoRL77yz775yx3HsVUWWeaMSaiMjR0IckjSUR\nQY7GF3clSaOLwS9JhTH4JakwBr8kFcbgl6TCGPySVBiDX5IKY/BLUmEMfkkqjMEvSYUx+CWpMAa/\nJBXG4Jekwhj8klQYg1+SCmPwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMJUCv6I+GZEvBwRTw9S\nc1VEbIyIJyPilCrjSZKqq3rFfxPQPdDOiDgXmJuZ84A/BK6pOJ4kjYgXXniB1atXs3379k630naV\ngj8zHwReHaTkU8C3GrWPANMj4rAqY0pSVVdd9Q3mzj2Rc85ZyjHHzOWuu+7qdEtt1bWfjz8L2NK0\nvhWYDby8n8eVpL3auHEjl1zyX3njja/wxhvTgS187nOfZ/v2nzB58uROt9cW+zv4AaLfeu6tqKen\nZ89yrVajVqvtv44kFWvTpk1MmDCL11+f3thyFDCBl156iWOPPbaTre1TvV6nXq9XPk5k7jWHh36A\niDnAdzPzpL3suxaoZ+aKxvp64MzMfLlfXVbtQ5KG4kc/+hEnnXQar7/+ReC3gOc4+ODvsG3bS2Pu\nij8iyMz+F9f7tL/fznkH8EWAiFgE/Lx/6EtSOx1//PH89V9fweTJN3HIITdy8MHfYeXKvx9zoV9F\npSv+iLgFOBOYSd+8/WXABIDMvK5RczV97/z5NfAHmfmDvRzHK35JbbVt2za2bt3Kcccdx6GHHtrp\ndlrS6hV/5amekWDwS9LwjdapHknSKGPwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/JJUGINf\nkgpj8EtSYQx+SSqMwS9JhTH4JakwBr8kFcbgl6TCGPySVBiDX5IKY/BLUmEMfkkqjMEvSYUx+CWp\nMAa/JBXG4JekwlQO/ojojoj1EbExIi7ey/5DI+K7EfFERPwwIr5UdUxJUusiM1t/cMR4YAPwceAF\nYA1wfmaua6r5U2BaZl4aETMb9Ydl5u6mmqzShySVKCLIzBju46pe8S8ENmXmc5m5C1gBLOlX0wsc\n0lg+BHilOfQlSe1VNfhnAVua1rc2tjW7GpgfES8CTwJfqzimJKmCroqPH8r8TDfwg8w8KyKOB1ZF\nxIcy87Xmop6enj3LtVqNWq1WsTVJOrDU63Xq9Xrl41Sd418E9GRmd2P9UqA3My9vqrkTWJaZDzXW\n7wcuzsy1TTXO8UvSMHVqjn8tMC8i5kTEROA84I5+Nc/T9+IvEXEY8H5gc8VxJUktqjTVk5m7I+Ii\n4B5gPHBjZq6LiAsb+68D/hz4u4h4CgjgTzLzZxX7liS1qNJUz4g14VSPJA1bp6Z6JEljjMEvSYUx\n+CWpMAa/JBXG4Jekwhj8klQYg1+SCmPwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/NIBbufO\nnWzYsIHduw+sP3Wdmbzyyivs3LlzSHW7du1qU2ejn8EvHcCWLVvGpElTOeGEk5g4cQrXX399p1sa\nEc8++yzz587lmCOPZPq0aVzzN3+z17qNGzfyvjlz9tTddNNNbe50dPLz+KUD1Lp165g//0PAUuAY\nYANwG9u2vcjMmTM721xFp37wg8xct45/3dvLq8DNU6Zwd73OggUL3lF34ty5HLt5M6dnsg1YftBB\n/L+HH+bkk0/uSN8jzc/jl/QO9957L/Bb9IU+9P3V08l8//vf71xTI6C3t5cn161jUW8vAcwA5vX2\nsmbNmnfUvfHGG2x69lkWNi4q/xUwd9w4Hnvssbb3PNoY/NIB6qSTTgJ+BvyqseVV4HVOPPHEzjU1\nAsaNG8d73/Menm+s7wZe6upi9uzZ76ibNGkS06ZOZWtjfSfwUsS76krkVI90ADvzzLP53vdWA7OA\nLfzu736aW29d0em2Klu1ahWf/fSnmdPVxba33mLxJz7BittvZ9y4d17L3nnnnSw97zzmdHXx8ltv\n8YlPfYpvLV9OxLBnR0alVqd6DH7pALd8+XLWrl1LrVZjyZIlnW5nxPz4xz/m0Ucf5b3vfS8f/ehH\nBwzzzZs389hjj3HEEUewePHiAyb0weCXpOL44q4kaUgMfkkqTOXgj4juiFgfERsj4uIBamoR8XhE\n/DAi6lXHlCS1rtIcf0SMp++ukI8DLwBrgPMzc11TzXTgIeCTmbk1ImZm5vZ+x3GOX5KGqVNz/AuB\nTZn5XGbuAlYA/d828Hng9szcCtA/9CVJ7VU1+PveHPwbWxvbms0DZkTEAxGxNiJ+r+KYkqQKuio+\nfijzMxOAU4GzgSnA6oh4ODM3Nhf19PTsWa7VatRqtYqtSdKBpV6vU6/XKx+n6hz/IqAnM7sb65cC\nvZl5eVPNxcBBmdnTWP9b4O7MvK2pxjl+SRqmTs3xrwXmRcSciJgInAfc0a/mH4EzImJ8REwBTgee\nqTiuJKlFlaZ6MnN3RFwE3AOMB27MzHURcWFj/3WZuT4i7gaeAnqBGzLT4JekDvEjGyRpjPIjGyRJ\nQ2LwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/JJUGINfkgpj8EtSYQx+SSqMwS9JhTH4Jakw\nBr8kFcbgl6TCGPySVBiDX5IKY/BLUmEMfkkqjMEvSYUx+CWpMAa/JBXG4JekwlQO/ojojoj1EbEx\nIi4epG5BROyOiM9UHVOS1LpKwR8R44GrgW5gPnB+RJw4QN3lwN1AVBlTklRN1Sv+hcCmzHwuM3cB\nK4Ale6n7KnAbsK3ieJKkiqoG/yxgS9P61sa2PSJiFn2/DK5pbMqKY0qSKuiq+PihhPiVwCWZmRER\nDDDV09PTs2e5VqtRq9UqtiZJB5Z6vU69Xq98nMhs/QI8IhYBPZnZ3Vi/FOjNzMubajbzm7CfCewA\n/n1m3tFUk1X6kKQSRQSZOezXTasGfxewATgbeBF4FDg/M9cNUH8T8N3MXNlvu8EvScPUavBXmurJ\nzN0RcRFwDzAeuDEz10XEhY3911U5viRp5FW64h+xJrzil6Rha/WK3zt3JakwBr8kFcbgl6TCGPwa\nc3bu3MkRRxxNxEQiJjJnzjzeeuutlo/32muv8dklSzh06lSOPvxwVq5cue8HSWOYL+5qzDnhhA+w\nYcN24N8BvcD/4vTT38fDD69u6XifXbKEzffcw8fefJNXgH846CDue/BBTjvttBHsWhp5vrirYmzc\n+GPg48AM+u4JPIvHHnum5ePds2oVZ7/5JgcDxwAf2LWL+++/f0R6lUYjg19jTlfXeOCVpi3bmTSp\n9VtSDp02bc/REvj5xIlMnz69QofS6Gbwa8y54or/BtwP3AH8A/AQ1177P1s+3v+4+mpWHnQQq7q6\nuG3KFMYddRRLly4doW6l0cc5fo1Jt956K5dddhnjx4/n8ssv59xzz610vDVr1nDfffcxY8YMli5d\nytSpU0eoU2n/6chn9YwUg1+Shs8XdyVJQ2LwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/JJU\nGINfkgpj8EtSYQx+SSqMwS9JhTH4JakwlYM/IrojYn1EbIyIi/ey/wsR8WREPBURD0XEyVXHlCS1\nrtLHMkfEeGADfX8H7wVgDXB+Zq5rqvkI8Exm/iIiuoGezFzU7zh+LLMkDVOnPpZ5IbApM5/LzF3A\nCmBJc0Fmrs7MXzRWHwFmVxxTklRB1eCfBWxpWt/a2DaQrwB3VRxTklRB63+hus+Q52ci4izgy8Di\nve3v6enZs1yr1ajVahVbk6QDS71ep16vVz5O1Tn+RfTN2Xc31i8FejPz8n51JwMrge7M3LSX4zjH\nL0nD1Kk5/rXAvIiYExETgfOAO/o1djR9ob90b6EvSWqvSlM9mbk7Ii4C7gHGAzdm5rqIuLCx/zrg\nz4D3ANdEBMCuzFxYrW1JUqsqTfWMWBNO9UjSsHVqqkeSNMYY/JJUGINfkgpj8EtSYQx+SSqMwS9J\nhTH4JakwBr8kFcbgl6TCGPySVBiDX5IKY/BLUmEMfkkqjMEvSYUx+CWpMAa/JBXG4Jekwhj8klQY\ng1+SCmPwS1JhDH5JKozBL0mFMfglqTCVgz8iuiNifURsjIiLB6i5qrH/yYg4peqYkqTWVQr+iBgP\nXA10A/OB8yPixH415wJzM3Me8IfANVXGHE1uv/12pk8/nIkTD+XDH17AL3/5y3fVZCZ/dcUVnDJ/\nPmcsWMD999+/Z/sVV/wV8+efwoIFZ+zZLkn7W2Rm6w+O+AhwWWZ2N9YvAcjMv2yquRZ4IDO/3Vhf\nD5yZmS831WSVPjphzZo1LFy4GPi3wGHAAxx11Dief37TO+r++1/8BdctW8bHduzgV8CqKVO494EH\nuPfe+1i27Dp27PgY8CumTFnFAw/cy8KFC9v/zUgakyKCzIzhPq7qVM8sYEvT+tbGtn3VzK44bsd9\n4xvfAE4AFgBHA59jy5Zn2b179zvqbrr+ej65YwdzgA8Cp+7YwS3Ll3PDDX/Hjh2fhMaeHTtOY/ny\nW9r6PUgqU1fFxw/1Mr3/b6R3Pa6np2fPcq1Wo1artdxUO0yePBl4o2nLG0AQ8c5vddKkSbzZtL5z\n3DgmT57MxIkToWnPuHFvMmnSpP3YsaSxrl6vU6/XKx+n6lTPIqCnaarnUqA3My9vqrkWqGfmisb6\nATHVs2XLFubMeR+9vR8EDgceYvHik/mnf6q/o+6WW27hqxdcwOk7dvDrCJ6eNo01jz/OI488wgUX\nfJUdO04n4tdMm/Y0jz++huOOO64T346kMajVqZ6qwd8FbADOBl4EHgXOz8x1TTXnAhdl5rmNXxRX\nZuaifscZc8EPsH79epYu/X1++tOf0d19Ftdeey3jxr179uzuu+/m2zffzNRp0/jjr3+duXPn7tl+\n883fZtq0qXz963+8Z7skDUVHgr8x8DnAlcB44MbMXBYRFwJk5nWNmrff+fNr4A8y8wf9jjEmg1+S\nOqljwT8SDH5JGr5OvatHkjTGGPySVBiDX5IKY/BLUmEMfkkqjMEvSYUx+CWpMAa/JBXG4Jekwhj8\nklQYg1+SCmPwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/JJUGINfkgpj8EtSYQx+SSqMwS9J\nhWk5+CNiRkSsioh/iYh7I2L6XmqOiogHIuKfI+KHEfFH1dqVJFVV5Yr/EmBVZr4PuL+x3t8u4D9l\n5geARcB/jIgTK4w5KtXr9U63UIn9d5b9d85Y7r2KKsH/KeBbjeVvAZ/uX5CZP8nMJxrLvwLWAUdW\nGHNUGutPHvvvLPvvnLHcexVVgv+wzHy5sfwycNhgxRExBzgFeKTCmJKkiroG2xkRq4DD97LrvzSv\nZGZGRA5ynIOB24CvNa78JUkdEpkD5vXgD4xYD9Qy8ycRcQTwQGaesJe6CcCdwP/JzCsHOFZrTUhS\n4TIzhvuYQa/49+EO4PeByxv//U7/gogI4EbgmYFCH1prXJLUmipX/DOAvweOBp4DPpeZP4+II4Eb\nMvO3I+IM4HvAU8DbA12amXdX7lyS1JKWg1+SNDZ15M7dsXrzV0R0R8T6iNgYERcPUHNVY/+TEXFK\nu3sczL76j4gvNPp+KiIeioiTO9HnQIby82/ULYiI3RHxmXb2N5ghPndqEfF44/leb3OLgxrCc+fQ\niPhuRDzR6P9LHWhzryLimxHxckQ8PUjNaD5vB+2/pfM2M9v+BVwB/Elj+WLgL/dSczjw4cbywcAG\n4MRO9NvoYTywCZgDTACe6N8PcC5wV2P5dODhTvXbYv8fAQ5tLHePtf6b6v4vfW8o+J1O9z2Mn/10\n4J+B2Y31mZ3ue5j9/ymw7O3egVeArk733ujn39D3VvKnB9g/as/bIfY/7PO2U5/VMxZv/loIbMrM\n5zJzF7ACWNKvZs/3lZmPANMjYtD7G9pon/1n5urM/EVj9RFgdpt7HMxQfv4AX6XvrcPb2tncPgyl\n988Dt2fmVoDM3N7mHgczlP57gUMay4cAr2Tm7jb2OKDMfBB4dZCS0Xze7rP/Vs7bTgX/WLz5axaw\npWl9a2PbvmpGS3gOpf9mXwHu2q8dDc8++4+IWfQF0jWNTaPlBayh/OznATMa05trI+L32tbdvg2l\n/6uB+RHxIvAk8LU29TYSRvN5O1xDOm+rvJ1zUAfgzV9DDZH+b00dLeEz5D4i4izgy8Di/dfOsA2l\n/yuBSxrPqeDd/y86ZSi9TwBOBc4GpgCrI+LhzNy4XzsbmqH03w38IDPPiojjgVUR8aHMfG0/9zZS\nRut5O2TDOW/3W/Bn5icG2td4oeLw/M3NXz8doG4CcDtwc2a+6z6BNnsBOKpp/Sj6rgwGq5nd2DYa\nDKV/Gi8M3QB0Z+Zg/zxut6H0fxqwoi/zmQmcExG7MvOO9rQ4oKH0vgXYnpmvA69HxPeADwGjIfiH\n0v+XgGUAmfmjiHgWeD+wth0NVjSaz9shGe5526mpnrdv/oKKN3+10VpgXkTMiYiJwHn0fR/N7gC+\nCBARi4CfN01pddo++4+Io4GVwNLM3NSBHgezz/4z87jMPDYzj6XvX4n/YRSEPgztufOPwBkRMT4i\nptD3IuMzbe5zIEPp/3ng4wCN+fH3A5vb2mXrRvN5u08tnbcdepV6BnAf8C/AvcD0xvYjgf/dWD6D\nvheMngAeb3x1d/jV9XPoe3fRJvpuRAO4ELiwqebqxv4ngVM72e9w+wf+lr53Y7z983600z0P9+ff\nVHsT8JlO9zzM585/pu+dPU8Df9Tpnof53DkCuIe+mzWfBj7f6Z6ber8FeBHYSd+/rL48xs7bQftv\n5bz1Bi5JKox/elGSCmPwS1JhDH5JKozBL0mFMfglqTAGvyQVxuCXpMIY/JJUmP8PQCcjVG1gpTcA\nAAAASUVORK5CYII=\n",
"text/plain": [
"<matplotlib.figure.Figure at 0x7fc98babab90>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"from matplotlib import pyplot as plt\n",
"%matplotlib inline\n",
"plt.scatter(humanize(predicted_y),humanize(actual_y),c=query_id)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Notes\n",
"\n",
"This implementation is __HIGHLY INEFFICIENT__ for multiple query ids. If ever use it, one is invited to split data into query IDs and use this less inefficient implementation.\n",
"\n",
"In fact, the one below is also \"bad\" because it computes all pairs of Qi,Qj while a more effective algorithm should only consider those of different actual_y. \n",
"Also, computing second derivatives the proposed way is inefficient as you compute the square matrix and that take diagonal.\n",
"\n",
"\n",
"But it's fucken' 2 AM and i wanna sleep, so sorry, i'm out."
]
},
{
"cell_type": "code",
"execution_count": 13,
"metadata": {
"collapsed": false
},
"outputs": [],
"source": [
"actual_y = T.vector()\n",
"predicted_y = T.vector()\n",
"\n",
"yi_lessthan_yj = actual_y[:,None] < actual_y[None,:]\n",
"loss_i_j = T.exp(predicted_y[:,None] - predicted_y[None,:])\n",
"\n",
"loss = T.sum(yi_lessthan_yj * loss_i_j)\n",
"\n",
"grad = T.grad(loss,predicted_y)\n",
"\n",
"hess = T.diagonal(T.hessian(loss,predicted_y))\n",
"\n",
"get_loss_and_derivatives_per_Qid = theano.function([predicted_y,actual_y],[loss,grad,hess])"
]
},
{
"cell_type": "code",
"execution_count": 14,
"metadata": {
"collapsed": false
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"10 loops, best of 3: 199 ms per loop\n"
]
}
],
"source": [
"%%timeit -n 10\n",
"actual_y = np.array([0,0,0,1,1,1,2,2,3,0,1,2,2]*20,'float32')\n",
"predicted_y = np.random.randn(len(actual_y))\n",
"\n",
"get_loss_and_derivatives_per_Qid(predicted_y,actual_y)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"collapsed": true
},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "Python [Root]",
"language": "python",
"name": "Python [Root]"
},
"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.12"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment