Skip to content

Instantly share code, notes, and snippets.

@ikegami-yukino
Created January 8, 2016 20:16
Show Gist options
  • Select an option

  • Save ikegami-yukino/01fc6121a46e390b4e1b to your computer and use it in GitHub Desktop.

Select an option

Save ikegami-yukino/01fc6121a46e390b4e1b to your computer and use it in GitHub Desktop.
scikit-learnのRandomForestのモデルをPythonコードに変換
Display the source blob
Display the rendered blob
Raw
{
"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