Created
January 8, 2016 20:16
-
-
Save ikegami-yukino/01fc6121a46e390b4e1b to your computer and use it in GitHub Desktop.
scikit-learnのRandomForestのモデルをPythonコードに変換
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": [ | |
| "# scikit-learnのRandomForestからPythonコード生成" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## 生成用コードをダウンロード" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 8, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "['--2016-01-09 05:08:24-- https://raw.githubusercontent.com/ikegami-yukino/misc/master/machinelearning/dt2code.py',\n", | |
| " 'Resolving raw.githubusercontent.com... 103.245.222.133',\n", | |
| " 'Connecting to raw.githubusercontent.com|103.245.222.133|:443... connected.',\n", | |
| " 'HTTP request sent, awaiting response... 200 OK',\n", | |
| " 'Length: 2052 (2.0K) [text/plain]',\n", | |
| " \"Saving to: 'dt2code.py'\",\n", | |
| " '',\n", | |
| " ' 0K .. 100% 122M=0s',\n", | |
| " '',\n", | |
| " \"2016-01-09 05:08:24 (122 MB/s) - 'dt2code.py' saved [2052/2052]\",\n", | |
| " '']" | |
| ] | |
| }, | |
| "execution_count": 8, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| } | |
| ], | |
| "source": [ | |
| "%system wget https://raw.githubusercontent.com/ikegami-yukino/misc/master/machinelearning/dt2code.py" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## irisを学習" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 1, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "from sklearn.datasets import load_iris\n", | |
| "from sklearn.ensemble import RandomForestClassifier\n", | |
| "\n", | |
| "iris = load_iris()\n", | |
| "clf = RandomForestClassifier(n_estimators=3)\n", | |
| "clf = clf.fit(iris.data, iris.target)" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## Random forestからコード生成\n", | |
| "\n", | |
| "`clf.estimators_`に決定木が複数入ってるのでループで回してそれぞれコード生成する" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 2, | |
| "metadata": { | |
| "collapsed": false, | |
| "scrolled": false | |
| }, | |
| "outputs": [ | |
| { | |
| "name": "stdout", | |
| "output_type": "stream", | |
| "text": [ | |
| "def func0(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if petal_width <= 0.800000011921: # samples=100\n", | |
| " return 0 # samples=33\n", | |
| " else:\n", | |
| " if petal_width <= 1.75: # samples=67\n", | |
| " if sepal_length <= 7.09999990463: # samples=40\n", | |
| " if petal_length <= 4.94999980927: # samples=39\n", | |
| " if sepal_length <= 4.94999980927: # samples=35\n", | |
| " if petal_length <= 3.90000009537: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=33\n", | |
| " else:\n", | |
| " if sepal_width <= 2.90000009537: # samples=4\n", | |
| " return 2 # samples=3\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if petal_length <= 4.85000038147: # samples=27\n", | |
| " if sepal_length <= 5.94999980927: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=25\n", | |
| "\n", | |
| "def func1(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if sepal_length <= 5.44999980927: # samples=95\n", | |
| " if petal_width <= 0.75: # samples=32\n", | |
| " return 0 # samples=28\n", | |
| " else:\n", | |
| " if sepal_length <= 4.94999980927: # samples=4\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=3\n", | |
| " else:\n", | |
| " if petal_length <= 4.85000038147: # samples=63\n", | |
| " if petal_width <= 0.600000023842: # samples=25\n", | |
| " return 0 # samples=1\n", | |
| " else:\n", | |
| " if sepal_width <= 2.95000004768: # samples=24\n", | |
| " return 1 # samples=14\n", | |
| " else:\n", | |
| " if petal_width <= 1.70000004768: # samples=10\n", | |
| " return 1 # samples=8\n", | |
| " else:\n", | |
| " if sepal_length <= 5.94999980927: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if petal_width <= 1.75: # samples=38\n", | |
| " if sepal_width <= 2.90000009537: # samples=6\n", | |
| " return 2 # samples=3\n", | |
| " else:\n", | |
| " if petal_length <= 5.40000009537: # samples=3\n", | |
| " return 1 # samples=2\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=32\n", | |
| "\n", | |
| "def func2(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if petal_width <= 0.800000011921: # samples=92\n", | |
| " return 0 # samples=29\n", | |
| " else:\n", | |
| " if petal_width <= 1.54999995232: # samples=63\n", | |
| " if petal_length <= 4.84999990463: # samples=28\n", | |
| " return 1 # samples=26\n", | |
| " else:\n", | |
| " return 2 # samples=2\n", | |
| " else:\n", | |
| " if petal_width <= 1.84999990463: # samples=35\n", | |
| " if petal_length <= 5.25: # samples=9\n", | |
| " if sepal_length <= 5.40000009537: # samples=5\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if sepal_width <= 3.09999990463: # samples=4\n", | |
| " if petal_length <= 4.94999980927: # samples=3\n", | |
| " return 2 # samples=2\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=4\n", | |
| " else:\n", | |
| " return 2 # samples=26\n", | |
| "\n" | |
| ] | |
| } | |
| ], | |
| "source": [ | |
| "import dt2code\n", | |
| "codes = [dt2code.dt2code(clf.estimators_[i], iris.feature_names, iris.target_names, 'func%s' % (i)) for i in range(len(clf.estimators_))]\n", | |
| "print('\\n'.join(codes))" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## できたコードを貼り付ける" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 3, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "def func0(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if petal_width <= 0.800000011921: # samples=100\n", | |
| " return 0 # samples=33\n", | |
| " else:\n", | |
| " if petal_width <= 1.75: # samples=67\n", | |
| " if sepal_length <= 7.09999990463: # samples=40\n", | |
| " if petal_length <= 4.94999980927: # samples=39\n", | |
| " if sepal_length <= 4.94999980927: # samples=35\n", | |
| " if petal_length <= 3.90000009537: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=33\n", | |
| " else:\n", | |
| " if sepal_width <= 2.90000009537: # samples=4\n", | |
| " return 2 # samples=3\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if petal_length <= 4.85000038147: # samples=27\n", | |
| " if sepal_length <= 5.94999980927: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=25\n", | |
| "\n", | |
| "def func1(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if sepal_length <= 5.44999980927: # samples=95\n", | |
| " if petal_width <= 0.75: # samples=32\n", | |
| " return 0 # samples=28\n", | |
| " else:\n", | |
| " if sepal_length <= 4.94999980927: # samples=4\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=3\n", | |
| " else:\n", | |
| " if petal_length <= 4.85000038147: # samples=63\n", | |
| " if petal_width <= 0.600000023842: # samples=25\n", | |
| " return 0 # samples=1\n", | |
| " else:\n", | |
| " if sepal_width <= 2.95000004768: # samples=24\n", | |
| " return 1 # samples=14\n", | |
| " else:\n", | |
| " if petal_width <= 1.70000004768: # samples=10\n", | |
| " return 1 # samples=8\n", | |
| " else:\n", | |
| " if sepal_length <= 5.94999980927: # samples=2\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if petal_width <= 1.75: # samples=38\n", | |
| " if sepal_width <= 2.90000009537: # samples=6\n", | |
| " return 2 # samples=3\n", | |
| " else:\n", | |
| " if petal_length <= 5.40000009537: # samples=3\n", | |
| " return 1 # samples=2\n", | |
| " else:\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=32\n", | |
| "\n", | |
| "def func2(sepal_length=0, sepal_width=0, petal_length=0, petal_width=0):\n", | |
| " \"\"\"\n", | |
| " 0 -> setosa\n", | |
| " 1 -> versicolor\n", | |
| " 2 -> virginica\n", | |
| " \"\"\"\n", | |
| " if petal_width <= 0.800000011921: # samples=92\n", | |
| " return 0 # samples=29\n", | |
| " else:\n", | |
| " if petal_width <= 1.54999995232: # samples=63\n", | |
| " if petal_length <= 4.84999990463: # samples=28\n", | |
| " return 1 # samples=26\n", | |
| " else:\n", | |
| " return 2 # samples=2\n", | |
| " else:\n", | |
| " if petal_width <= 1.84999990463: # samples=35\n", | |
| " if petal_length <= 5.25: # samples=9\n", | |
| " if sepal_length <= 5.40000009537: # samples=5\n", | |
| " return 2 # samples=1\n", | |
| " else:\n", | |
| " if sepal_width <= 3.09999990463: # samples=4\n", | |
| " if petal_length <= 4.94999980927: # samples=3\n", | |
| " return 2 # samples=2\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 1 # samples=1\n", | |
| " else:\n", | |
| " return 2 # samples=4\n", | |
| " else:\n", | |
| " return 2 # samples=26" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## 識別関数を作る\n", | |
| "\n", | |
| "それぞれの木が出力したクラスから多数決で決める" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 4, | |
| "metadata": { | |
| "collapsed": true | |
| }, | |
| "outputs": [], | |
| "source": [ | |
| "from collections import Counter\n", | |
| "\n", | |
| "\n", | |
| "def predict(X):\n", | |
| " def majority_decision(*args):\n", | |
| " FUNCS = (func0, func1, func2)\n", | |
| " cntr = Counter()\n", | |
| " for f in FUNCS:\n", | |
| " predicted = f(*args)\n", | |
| " cntr[predicted] += 1\n", | |
| " return cntr.most_common()[0][0]\n", | |
| " return list(map(lambda x: majority_decision(*x), X))" | |
| ] | |
| }, | |
| { | |
| "cell_type": "markdown", | |
| "metadata": {}, | |
| "source": [ | |
| "## 結果\n", | |
| "sklearnと同じ結果が出る" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 5, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "[0, 1, 2]" | |
| ] | |
| }, | |
| "execution_count": 5, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| } | |
| ], | |
| "source": [ | |
| "predict([[5.1, 3.5, 1.4, 0.2], [5, 2, 3.5, 1], [6.3, 3.3, 6, 2.5]])" | |
| ] | |
| }, | |
| { | |
| "cell_type": "code", | |
| "execution_count": 6, | |
| "metadata": { | |
| "collapsed": false | |
| }, | |
| "outputs": [ | |
| { | |
| "data": { | |
| "text/plain": [ | |
| "array([0, 1, 2])" | |
| ] | |
| }, | |
| "execution_count": 6, | |
| "metadata": {}, | |
| "output_type": "execute_result" | |
| } | |
| ], | |
| "source": [ | |
| "clf.predict([[5.1, 3.5, 1.4, 0.2], [5, 2, 3.5, 1], [6.3, 3.3, 6, 2.5]])" | |
| ] | |
| }, | |
| { | |
| "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.10" | |
| } | |
| }, | |
| "nbformat": 4, | |
| "nbformat_minor": 0 | |
| } |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment