Last active
October 28, 2016 16:54
-
-
Save justheuristic/177467f866cd4f8b4a7b8150a47a5953 to your computer and use it in GitHub Desktop.
Parallel minibatch iterator for tensorflow and theano
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
| """ | |
| ## What is it? | |
| # mostly from http://stackoverflow.com/questions/7323664/python-generator-pre-fetch | |
| The only parameter here is | |
| """ | |
| ###background generator (from your generator) | |
| import threading | |
| import sys | |
| if sys.version_info >= (3, 0): | |
| import queue as Queue | |
| else: | |
| import Queue | |
| class BackgroundGenerator(threading.Thread): | |
| def __init__(self, generator, max_prefetch=1): | |
| """ | |
| This function transforms generator into a background-thead generator. | |
| :param generator: generator or genexp or any | |
| It can be used with any minibatch generator. | |
| It is quite lightweight, but not entirely weightless. | |
| Using global variables inside generator is not recommended (may rise GIL and zero-out the benefit of having a background thread.) | |
| The ideal use case is when everything it requires is store inside it and everything it outputs is passed through queue. | |
| There's no restriction on reading/writing files or retrieving URLs [or whatever] wlilst iterating. | |
| :param max_prefetch: defines, how many iterations (at most) can background generator keep stored at any moment of time. | |
| Whenever there's already max_prefetch batches stored in queue, the background process will halt until one of these batches is dequeued. | |
| !Default max_prefetch=1 is okay unless you deal with some weird file IO in your generator! | |
| Setting max_prefetch to -1 lets it store as many batches as it can, which will work slightly (if any) faster, but will require storing | |
| all batches in memory. If you use infinite generator with max_prefetch=-1, it will exceed the RAM size unless dequeued quickly enough. | |
| """ | |
| threading.Thread.__init__(self) | |
| self.queue = Queue.Queue(max_prefetch) | |
| self.generator = generator | |
| self.daemon = True | |
| self.start() | |
| def run(self): | |
| for item in self.generator: | |
| self.queue.put(item) | |
| self.queue.put(None) | |
| def next(self): | |
| next_item = self.queue.get() | |
| if next_item is None: | |
| raise StopIteration | |
| return next_item | |
| # Python 3 compatibility | |
| def __next__(self): | |
| return self.next() | |
| def __iter__(self): | |
| return self | |
| #decorator | |
| class background: | |
| def __init__(self,max_prefetch=1): | |
| self.max_prefetch = max_prefetch | |
| def __call__(self,gen): | |
| def bg_generator(*args,**kwargs): | |
| return BackgroundGenerator(gen(*args,**kwargs)) | |
| return bg_generator | |
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": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## Что это такое\n", | |
| "\n", | |
| "Это штука, которая делает из генератора фоновый генератор. \n", | |
| "Его можно запихнуть в модуль и использовать с любым итератором минибатчей. \n", | |
| "Он использует только threading и Queue (питонячие стандартные библиотеки).\n", | |
| "\n", | |
| "Он довольно лёгкий, но не абсолютно невесомый.\n", | |
| "Лучше не использовать в нём переменные из \"внешнего мира\" - общаться с внешним миром только через итерацию.\n", | |
| "Хороший способ выстрелить себе в ногу - дёргать им GIL.\n", | |
| "\n", | |
| "\n", | |
| "## Про тюнинг\n", | |
| "\n", | |
| "\n", | |
| "__Если у вас простой итератор без извращений - не читайте это__\n", | |
| "\n", | |
| "\n", | |
| "Генератор предзапасает не более max_prefetch батчей. На практике менять этот параметр нужно только если ваш генератор неравномерен по времени на итерацию. Например, раз в 100 итераций, читает новый файл с диска.\n", | |
| "\n", | |
| "Чем больше max_prefetch, тем меньше оверхеда по времени у генератора, но тем больше батчей нужно хранить в памяти. Если max_prefetch = -1, генератор будет запасать столько батчей, сколько успеет. Это работает чуть быстрее, чем обычно, но на бесконечном генераторе приведёт к переполнению памяти.\n", | |
| "\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 1, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "#mostly from http://stackoverflow.com/questions/7323664/python-generator-pre-fetch\n", | |
| "\n", | |
| "\n", | |
| "###background generator (from your generator)\n", | |
| "import threading\n", | |
| "import Queue\n", | |
| "\n", | |
| "class BackgroundGenerator(threading.Thread):\n", | |
| " def __init__(self, generator,max_prefetch = 1):\n", | |
| " threading.Thread.__init__(self)\n", | |
| " self.queue = Queue.Queue(max_prefetch)\n", | |
| " self.generator = generator\n", | |
| " self.daemon = True\n", | |
| " self.start()\n", | |
| "\n", | |
| " def run(self):\n", | |
| " for item in self.generator:\n", | |
| " self.queue.put(item)\n", | |
| " self.queue.put(None)\n", | |
| "\n", | |
| " def next(self):\n", | |
| " next_item = self.queue.get()\n", | |
| " if next_item is None:\n", | |
| " raise StopIteration\n", | |
| " return next_item\n", | |
| " \n", | |
| " # Python 3 compatibility\n", | |
| " def __next__(self):\n", | |
| " return self.next()\n", | |
| " \n", | |
| " def __iter__(self):\n", | |
| " return self\n", | |
| " \n", | |
| " " | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 2, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "###your super-mega data iterator\n", | |
| "import numpy as np\n", | |
| "import time\n", | |
| "\n", | |
| "def iterate_minibatches(n_batches, batch_size=10):\n", | |
| " for b_i in range(n_batches):\n", | |
| " time.sleep(0.1)\n", | |
| " X = np.random.normal(size=[batch_size,20])\n", | |
| " y = np.random.randint(0,2,size=batch_size)\n", | |
| " yield X,y" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### без параллелизма" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 3, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "name": "stdout", | |
| "output_type": "stream", | |
| "text": [ | |
| "! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! !\n", | |
| "CPU times: user 112 ms, sys: 12 ms, total: 124 ms\n", | |
| "Wall time: 10.1 s\n" | |
| ] | |
| } | |
| ], | |
| "source": [ | |
| "%%time\n", | |
| "for b_x,b_y in iterate_minibatches(50):\n", | |
| " #training\n", | |
| " time.sleep(0.1)\n", | |
| " print '!',\n", | |
| "print" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### с параллелизмом" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 4, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "name": "stdout", | |
| "output_type": "stream", | |
| "text": [ | |
| "! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! !\n", | |
| "CPU times: user 71.1 ms, sys: 8.59 ms, total: 79.7 ms\n", | |
| "Wall time: 5.14 s\n" | |
| ] | |
| } | |
| ], | |
| "source": [ | |
| "%%time\n", | |
| "for b_x,b_y in BackgroundGenerator(iterate_minibatches(50)):\n", | |
| " #training\n", | |
| " time.sleep(0.1)\n", | |
| " print '!',\n", | |
| "print" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "### Для любителей декораторов\n", | |
| "Если вам не нравится явный вызов BackgroundGenerator, можно сделать так" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 5, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "#decorator\n", | |
| "def background(gen):\n", | |
| " def bg_generator(*args,**kwargs):\n", | |
| " return BackgroundGenerator(gen(*args,**kwargs))\n", | |
| " \n", | |
| " return bg_generator\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 6, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "import numpy as np\n", | |
| "@background\n", | |
| "def bg_iterate_minibatches(n_batches, batch_size=10):\n", | |
| " for b_i in range(n_batches):\n", | |
| " X = np.random.normal(size=[batch_size,20])\n", | |
| " y = np.random.randint(0,2,size=batch_size)\n", | |
| " yield X,y\n", | |
| "\n" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 7, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "name": "stdout", | |
| "output_type": "stream", | |
| "text": [ | |
| "! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! ! !\n", | |
| "CPU times: user 61.2 ms, sys: 19.9 ms, total: 81.1 ms\n", | |
| "Wall time: 5.04 s\n" | |
| ] | |
| } | |
| ], | |
| "source": [ | |
| "%%time\n", | |
| "#параллелизм добавлен в декораторе\n", | |
| "for b_x,b_y in bg_iterate_minibatches(50):\n", | |
| " #training\n", | |
| " time.sleep(0.1)\n", | |
| " print '!',\n", | |
| "print" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": null, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [] | |
| } | |
| ], | |
| "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