Initial commit
This commit is contained in:
commit
8542a1f9ef
59 files changed
+73464
No files matched your search
Vendored
+12
@@ -0,0 +1,12 @@
|
||||
{
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "Python Debugger: Current File",
|
||||
"type": "debugpy",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal"
|
||||
}
|
||||
]
|
||||
}
|
||||
Vendored
+3
@@ -0,0 +1,3 @@
|
||||
{
|
||||
"terminal.integrated.defaultProfile.windows": "Powershell (conda)"
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,687 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "73b96d3d-406a-405d-850b-38dd7772624e",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# 决策树\n",
|
||||
"\n",
|
||||
"## 目录\n",
|
||||
"\n",
|
||||
"- 决策树学习器"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "72a368ee-02cc-4d92-8d0f-d4323d31d12c",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"source": [
|
||||
"## 决策树学习器\n",
|
||||
"\n",
|
||||
"### 概述\n",
|
||||
"\n",
|
||||
"#### 决策树\n",
|
||||
"\n",
|
||||
"决策树是一个流程图,它使用决策树及其可能的后果进行分类。在树的每个非叶子节点上,输入的一个属性被测试,根据这个测试结果,选择通往子节点的相应分支。在叶子节点上,根据这个叶子节点的类别标签,对输入进行分类。从根到叶的路径代表分类规则,根据这些规则给叶节点分配类标签。\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"#### 决策树学习\n",
|
||||
"\n",
|
||||
"决策树学习是指从有类标记的训练数据中构建决策树。数据预计是一个元组,其中元组的每个记录都是用于分类的属性。决策树是自上而下构建的,通过在每一步选择一个变量来最好地分割项目集。有不同的指标来衡量 \"最佳分割\"。这些指标通常衡量子集内目标变量的同质性。\n",
|
||||
"\n",
|
||||
"#### 信息增益\n",
|
||||
"\n",
|
||||
"信息增益是基于信息理论中的熵的概念。熵的定义为:\n",
|
||||
"\n",
|
||||
"$$H(p) = -\\sum{p_i \\log_2{p_i}}$$\n",
|
||||
"\n",
|
||||
"信息增益是指父代的熵和子代的熵的加权和之间的差异。用于分割的特征是提供最大信息增益的特征。\n",
|
||||
"\n",
|
||||
"#### 伪代码"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "1744317d-b45d-440c-b3ed-77ab64b6e10e",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"__function__ DECISION-TREE-LEARNING(_examples_, _attributes_, _parent\\_examples_) __returns__ a tree \n",
|
||||
" __if__ _examples_ 是空集 __then return__ PLURALITY\\-VALUE(_parent\\_examples_) \n",
|
||||
" __else if__ _examples_ 分类结果都相同 __then return__ 分类结果 \n",
|
||||
" __else if__ _attributes_ 是空集 __then return__ PLURALITY\\-VALUE(_examples_) \n",
|
||||
" __else__ \n",
|
||||
"   _A_ ← argmax<sub>_a_ ∈ _attributes_</sub> IMPORTANCE(_a_, _examples_) \n",
|
||||
"   _tree_ ← 以特征 _A_ 为根检测节点的决策树 \n",
|
||||
"   __for each__ 特征 _A_ 的取值 _v<sub>k</sub>_ __do__ \n",
|
||||
"     _exs_ ← \\{ _e_ : _e_ ∈ _examples_ __and__ _e_._A_ = _v<sub>k</sub>_ \\} \n",
|
||||
"     _subtree_ ← DECISION-TREE-LEARNING(_exs_, _attributes_ − _A_, _examples_) \n",
|
||||
"     将标签为 \\(_A_ = _v<sub>k</sub>_\\) 的子树 _subtree_ 添加为 _tree_ 的分支 \n",
|
||||
"   __return__ _tree_ "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "f8565216-f7c5-4b7f-a791-ebf538a6aaab",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 实现\n",
|
||||
"\n",
|
||||
"由我们的学习算法构建的树的节点,根据它们是内部节点还是叶节点,分别使用`DecisionFork`或`DecisionLeaf`来存储。"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"id": "13283545-a69a-4d96-ab5d-157eaa395f3a",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import sys\n",
|
||||
"sys.path.insert(1, '../')\n",
|
||||
"from utils.utils import *\n",
|
||||
"\n",
|
||||
"from utils.dataset4learners import *\n",
|
||||
"from DecisionTreeLearner_3 import *"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"id": "c3e70aae-3ef5-4e41-b39d-bc0f0b57270a",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<!DOCTYPE html PUBLIC \"-//W3C//DTD HTML 4.01//EN\"\n",
|
||||
" \"http://www.w3.org/TR/html4/strict.dtd\">\n",
|
||||
"<!--\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"-->\n",
|
||||
"<html>\n",
|
||||
"<head>\n",
|
||||
" <title></title>\n",
|
||||
" <meta http-equiv=\"content-type\" content=\"text/html; charset=None\">\n",
|
||||
" <style type=\"text/css\">\n",
|
||||
"/*\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"*/\n",
|
||||
"pre { line-height: 125%; }\n",
|
||||
"td.linenos .normal { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"td.linenos .special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos.special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"body .hll { background-color: #ffffcc }\n",
|
||||
"body { background: #f8f8f8; }\n",
|
||||
"body .c { color: #3D7B7B; font-style: italic } /* Comment */\n",
|
||||
"body .err { border: 1px solid #FF0000 } /* Error */\n",
|
||||
"body .k { color: #008000; font-weight: bold } /* Keyword */\n",
|
||||
"body .o { color: #666666 } /* Operator */\n",
|
||||
"body .ch { color: #3D7B7B; font-style: italic } /* Comment.Hashbang */\n",
|
||||
"body .cm { color: #3D7B7B; font-style: italic } /* Comment.Multiline */\n",
|
||||
"body .cp { color: #9C6500 } /* Comment.Preproc */\n",
|
||||
"body .cpf { color: #3D7B7B; font-style: italic } /* Comment.PreprocFile */\n",
|
||||
"body .c1 { color: #3D7B7B; font-style: italic } /* Comment.Single */\n",
|
||||
"body .cs { color: #3D7B7B; font-style: italic } /* Comment.Special */\n",
|
||||
"body .gd { color: #A00000 } /* Generic.Deleted */\n",
|
||||
"body .ge { font-style: italic } /* Generic.Emph */\n",
|
||||
"body .gr { color: #E40000 } /* Generic.Error */\n",
|
||||
"body .gh { color: #000080; font-weight: bold } /* Generic.Heading */\n",
|
||||
"body .gi { color: #008400 } /* Generic.Inserted */\n",
|
||||
"body .go { color: #717171 } /* Generic.Output */\n",
|
||||
"body .gp { color: #000080; font-weight: bold } /* Generic.Prompt */\n",
|
||||
"body .gs { font-weight: bold } /* Generic.Strong */\n",
|
||||
"body .gu { color: #800080; font-weight: bold } /* Generic.Subheading */\n",
|
||||
"body .gt { color: #0044DD } /* Generic.Traceback */\n",
|
||||
"body .kc { color: #008000; font-weight: bold } /* Keyword.Constant */\n",
|
||||
"body .kd { color: #008000; font-weight: bold } /* Keyword.Declaration */\n",
|
||||
"body .kn { color: #008000; font-weight: bold } /* Keyword.Namespace */\n",
|
||||
"body .kp { color: #008000 } /* Keyword.Pseudo */\n",
|
||||
"body .kr { color: #008000; font-weight: bold } /* Keyword.Reserved */\n",
|
||||
"body .kt { color: #B00040 } /* Keyword.Type */\n",
|
||||
"body .m { color: #666666 } /* Literal.Number */\n",
|
||||
"body .s { color: #BA2121 } /* Literal.String */\n",
|
||||
"body .na { color: #687822 } /* Name.Attribute */\n",
|
||||
"body .nb { color: #008000 } /* Name.Builtin */\n",
|
||||
"body .nc { color: #0000FF; font-weight: bold } /* Name.Class */\n",
|
||||
"body .no { color: #880000 } /* Name.Constant */\n",
|
||||
"body .nd { color: #AA22FF } /* Name.Decorator */\n",
|
||||
"body .ni { color: #717171; font-weight: bold } /* Name.Entity */\n",
|
||||
"body .ne { color: #CB3F38; font-weight: bold } /* Name.Exception */\n",
|
||||
"body .nf { color: #0000FF } /* Name.Function */\n",
|
||||
"body .nl { color: #767600 } /* Name.Label */\n",
|
||||
"body .nn { color: #0000FF; font-weight: bold } /* Name.Namespace */\n",
|
||||
"body .nt { color: #008000; font-weight: bold } /* Name.Tag */\n",
|
||||
"body .nv { color: #19177C } /* Name.Variable */\n",
|
||||
"body .ow { color: #AA22FF; font-weight: bold } /* Operator.Word */\n",
|
||||
"body .w { color: #bbbbbb } /* Text.Whitespace */\n",
|
||||
"body .mb { color: #666666 } /* Literal.Number.Bin */\n",
|
||||
"body .mf { color: #666666 } /* Literal.Number.Float */\n",
|
||||
"body .mh { color: #666666 } /* Literal.Number.Hex */\n",
|
||||
"body .mi { color: #666666 } /* Literal.Number.Integer */\n",
|
||||
"body .mo { color: #666666 } /* Literal.Number.Oct */\n",
|
||||
"body .sa { color: #BA2121 } /* Literal.String.Affix */\n",
|
||||
"body .sb { color: #BA2121 } /* Literal.String.Backtick */\n",
|
||||
"body .sc { color: #BA2121 } /* Literal.String.Char */\n",
|
||||
"body .dl { color: #BA2121 } /* Literal.String.Delimiter */\n",
|
||||
"body .sd { color: #BA2121; font-style: italic } /* Literal.String.Doc */\n",
|
||||
"body .s2 { color: #BA2121 } /* Literal.String.Double */\n",
|
||||
"body .se { color: #AA5D1F; font-weight: bold } /* Literal.String.Escape */\n",
|
||||
"body .sh { color: #BA2121 } /* Literal.String.Heredoc */\n",
|
||||
"body .si { color: #A45A77; font-weight: bold } /* Literal.String.Interpol */\n",
|
||||
"body .sx { color: #008000 } /* Literal.String.Other */\n",
|
||||
"body .sr { color: #A45A77 } /* Literal.String.Regex */\n",
|
||||
"body .s1 { color: #BA2121 } /* Literal.String.Single */\n",
|
||||
"body .ss { color: #19177C } /* Literal.String.Symbol */\n",
|
||||
"body .bp { color: #008000 } /* Name.Builtin.Pseudo */\n",
|
||||
"body .fm { color: #0000FF } /* Name.Function.Magic */\n",
|
||||
"body .vc { color: #19177C } /* Name.Variable.Class */\n",
|
||||
"body .vg { color: #19177C } /* Name.Variable.Global */\n",
|
||||
"body .vi { color: #19177C } /* Name.Variable.Instance */\n",
|
||||
"body .vm { color: #19177C } /* Name.Variable.Magic */\n",
|
||||
"body .il { color: #666666 } /* Literal.Number.Integer.Long */\n",
|
||||
"\n",
|
||||
" </style>\n",
|
||||
"</head>\n",
|
||||
"<body>\n",
|
||||
"<h2></h2>\n",
|
||||
"\n",
|
||||
"<div class=\"highlight\"><pre><span></span><span class=\"k\">class</span> <span class=\"nc\">DecisionFork</span><span class=\"p\">:</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""</span>\n",
|
||||
"<span class=\"sd\"> A fork of a decision tree holds an attribute to test, and a dict</span>\n",
|
||||
"<span class=\"sd\"> of branches, one for each of the attribute's values.</span>\n",
|
||||
"<span class=\"sd\"> """</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__init__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">attr</span><span class=\"p\">,</span> <span class=\"n\">attr_name</span><span class=\"o\">=</span><span class=\"kc\">None</span><span class=\"p\">,</span> <span class=\"n\">default_child</span><span class=\"o\">=</span><span class=\"kc\">None</span><span class=\"p\">,</span> <span class=\"n\">branches</span><span class=\"o\">=</span><span class=\"kc\">None</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Initialize by saying what attribute this node tests."""</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr</span> <span class=\"o\">=</span> <span class=\"n\">attr</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr_name</span> <span class=\"o\">=</span> <span class=\"n\">attr_name</span> <span class=\"ow\">or</span> <span class=\"n\">attr</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">default_child</span> <span class=\"o\">=</span> <span class=\"n\">default_child</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span> <span class=\"o\">=</span> <span class=\"n\">branches</span> <span class=\"ow\">or</span> <span class=\"p\">{}</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__call__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">example</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Given an example, classify it using the attribute and the branches."""</span>\n",
|
||||
" <span class=\"n\">attr_val</span> <span class=\"o\">=</span> <span class=\"n\">example</span><span class=\"p\">[</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr</span><span class=\"p\">]</span>\n",
|
||||
" <span class=\"k\">if</span> <span class=\"n\">attr_val</span> <span class=\"ow\">in</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span><span class=\"p\">:</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span><span class=\"p\">[</span><span class=\"n\">attr_val</span><span class=\"p\">](</span><span class=\"n\">example</span><span class=\"p\">)</span>\n",
|
||||
" <span class=\"k\">else</span><span class=\"p\">:</span>\n",
|
||||
" <span class=\"c1\"># return default class when attribute is unknown</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">default_child</span><span class=\"p\">(</span><span class=\"n\">example</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">add</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">val</span><span class=\"p\">,</span> <span class=\"n\">subtree</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Add a branch. If self.attr = val, go to the given subtree."""</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span><span class=\"p\">[</span><span class=\"n\">val</span><span class=\"p\">]</span> <span class=\"o\">=</span> <span class=\"n\">subtree</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">display</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">indent</span><span class=\"o\">=</span><span class=\"mi\">0</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"n\">name</span> <span class=\"o\">=</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr_name</span>\n",
|
||||
" <span class=\"nb\">print</span><span class=\"p\">(</span><span class=\"s1\">'Test'</span><span class=\"p\">,</span> <span class=\"n\">name</span><span class=\"p\">)</span>\n",
|
||||
" <span class=\"k\">for</span> <span class=\"p\">(</span><span class=\"n\">val</span><span class=\"p\">,</span> <span class=\"n\">subtree</span><span class=\"p\">)</span> <span class=\"ow\">in</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span><span class=\"o\">.</span><span class=\"n\">items</span><span class=\"p\">():</span>\n",
|
||||
" <span class=\"nb\">print</span><span class=\"p\">(</span><span class=\"s1\">' '</span> <span class=\"o\">*</span> <span class=\"mi\">4</span> <span class=\"o\">*</span> <span class=\"n\">indent</span><span class=\"p\">,</span> <span class=\"n\">name</span><span class=\"p\">,</span> <span class=\"s1\">'='</span><span class=\"p\">,</span> <span class=\"n\">val</span><span class=\"p\">,</span> <span class=\"s1\">'==>'</span><span class=\"p\">,</span> <span class=\"n\">end</span><span class=\"o\">=</span><span class=\"s1\">' '</span><span class=\"p\">)</span>\n",
|
||||
" <span class=\"n\">subtree</span><span class=\"o\">.</span><span class=\"n\">display</span><span class=\"p\">(</span><span class=\"n\">indent</span> <span class=\"o\">+</span> <span class=\"mi\">1</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__repr__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"s1\">'DecisionFork(</span><span class=\"si\">{0!r}</span><span class=\"s1\">, </span><span class=\"si\">{1!r}</span><span class=\"s1\">, </span><span class=\"si\">{2!r}</span><span class=\"s1\">)'</span><span class=\"o\">.</span><span class=\"n\">format</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr</span><span class=\"p\">,</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">attr_name</span><span class=\"p\">,</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">branches</span><span class=\"p\">)</span>\n",
|
||||
"</pre></div>\n",
|
||||
"</body>\n",
|
||||
"</html>\n"
|
||||
],
|
||||
"text/plain": [
|
||||
"<IPython.core.display.HTML object>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"psource(DecisionFork)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0ef40288-b753-4155-a813-25153163044f",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"`DecisionFork`持有属性,在该节点进行测试,以及一个分支的决定。分支存储了子节点,每个属性的值都有一个。以输入元组为参数,以函数形式调用这个类的对象,根据属性测试的结果返回分类路径中的下一个节点。"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"id": "79120735-ed3a-4b36-9fb2-5b590de076d3",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<!DOCTYPE html PUBLIC \"-//W3C//DTD HTML 4.01//EN\"\n",
|
||||
" \"http://www.w3.org/TR/html4/strict.dtd\">\n",
|
||||
"<!--\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"-->\n",
|
||||
"<html>\n",
|
||||
"<head>\n",
|
||||
" <title></title>\n",
|
||||
" <meta http-equiv=\"content-type\" content=\"text/html; charset=None\">\n",
|
||||
" <style type=\"text/css\">\n",
|
||||
"/*\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"*/\n",
|
||||
"pre { line-height: 125%; }\n",
|
||||
"td.linenos .normal { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"td.linenos .special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos.special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"body .hll { background-color: #ffffcc }\n",
|
||||
"body { background: #f8f8f8; }\n",
|
||||
"body .c { color: #3D7B7B; font-style: italic } /* Comment */\n",
|
||||
"body .err { border: 1px solid #FF0000 } /* Error */\n",
|
||||
"body .k { color: #008000; font-weight: bold } /* Keyword */\n",
|
||||
"body .o { color: #666666 } /* Operator */\n",
|
||||
"body .ch { color: #3D7B7B; font-style: italic } /* Comment.Hashbang */\n",
|
||||
"body .cm { color: #3D7B7B; font-style: italic } /* Comment.Multiline */\n",
|
||||
"body .cp { color: #9C6500 } /* Comment.Preproc */\n",
|
||||
"body .cpf { color: #3D7B7B; font-style: italic } /* Comment.PreprocFile */\n",
|
||||
"body .c1 { color: #3D7B7B; font-style: italic } /* Comment.Single */\n",
|
||||
"body .cs { color: #3D7B7B; font-style: italic } /* Comment.Special */\n",
|
||||
"body .gd { color: #A00000 } /* Generic.Deleted */\n",
|
||||
"body .ge { font-style: italic } /* Generic.Emph */\n",
|
||||
"body .gr { color: #E40000 } /* Generic.Error */\n",
|
||||
"body .gh { color: #000080; font-weight: bold } /* Generic.Heading */\n",
|
||||
"body .gi { color: #008400 } /* Generic.Inserted */\n",
|
||||
"body .go { color: #717171 } /* Generic.Output */\n",
|
||||
"body .gp { color: #000080; font-weight: bold } /* Generic.Prompt */\n",
|
||||
"body .gs { font-weight: bold } /* Generic.Strong */\n",
|
||||
"body .gu { color: #800080; font-weight: bold } /* Generic.Subheading */\n",
|
||||
"body .gt { color: #0044DD } /* Generic.Traceback */\n",
|
||||
"body .kc { color: #008000; font-weight: bold } /* Keyword.Constant */\n",
|
||||
"body .kd { color: #008000; font-weight: bold } /* Keyword.Declaration */\n",
|
||||
"body .kn { color: #008000; font-weight: bold } /* Keyword.Namespace */\n",
|
||||
"body .kp { color: #008000 } /* Keyword.Pseudo */\n",
|
||||
"body .kr { color: #008000; font-weight: bold } /* Keyword.Reserved */\n",
|
||||
"body .kt { color: #B00040 } /* Keyword.Type */\n",
|
||||
"body .m { color: #666666 } /* Literal.Number */\n",
|
||||
"body .s { color: #BA2121 } /* Literal.String */\n",
|
||||
"body .na { color: #687822 } /* Name.Attribute */\n",
|
||||
"body .nb { color: #008000 } /* Name.Builtin */\n",
|
||||
"body .nc { color: #0000FF; font-weight: bold } /* Name.Class */\n",
|
||||
"body .no { color: #880000 } /* Name.Constant */\n",
|
||||
"body .nd { color: #AA22FF } /* Name.Decorator */\n",
|
||||
"body .ni { color: #717171; font-weight: bold } /* Name.Entity */\n",
|
||||
"body .ne { color: #CB3F38; font-weight: bold } /* Name.Exception */\n",
|
||||
"body .nf { color: #0000FF } /* Name.Function */\n",
|
||||
"body .nl { color: #767600 } /* Name.Label */\n",
|
||||
"body .nn { color: #0000FF; font-weight: bold } /* Name.Namespace */\n",
|
||||
"body .nt { color: #008000; font-weight: bold } /* Name.Tag */\n",
|
||||
"body .nv { color: #19177C } /* Name.Variable */\n",
|
||||
"body .ow { color: #AA22FF; font-weight: bold } /* Operator.Word */\n",
|
||||
"body .w { color: #bbbbbb } /* Text.Whitespace */\n",
|
||||
"body .mb { color: #666666 } /* Literal.Number.Bin */\n",
|
||||
"body .mf { color: #666666 } /* Literal.Number.Float */\n",
|
||||
"body .mh { color: #666666 } /* Literal.Number.Hex */\n",
|
||||
"body .mi { color: #666666 } /* Literal.Number.Integer */\n",
|
||||
"body .mo { color: #666666 } /* Literal.Number.Oct */\n",
|
||||
"body .sa { color: #BA2121 } /* Literal.String.Affix */\n",
|
||||
"body .sb { color: #BA2121 } /* Literal.String.Backtick */\n",
|
||||
"body .sc { color: #BA2121 } /* Literal.String.Char */\n",
|
||||
"body .dl { color: #BA2121 } /* Literal.String.Delimiter */\n",
|
||||
"body .sd { color: #BA2121; font-style: italic } /* Literal.String.Doc */\n",
|
||||
"body .s2 { color: #BA2121 } /* Literal.String.Double */\n",
|
||||
"body .se { color: #AA5D1F; font-weight: bold } /* Literal.String.Escape */\n",
|
||||
"body .sh { color: #BA2121 } /* Literal.String.Heredoc */\n",
|
||||
"body .si { color: #A45A77; font-weight: bold } /* Literal.String.Interpol */\n",
|
||||
"body .sx { color: #008000 } /* Literal.String.Other */\n",
|
||||
"body .sr { color: #A45A77 } /* Literal.String.Regex */\n",
|
||||
"body .s1 { color: #BA2121 } /* Literal.String.Single */\n",
|
||||
"body .ss { color: #19177C } /* Literal.String.Symbol */\n",
|
||||
"body .bp { color: #008000 } /* Name.Builtin.Pseudo */\n",
|
||||
"body .fm { color: #0000FF } /* Name.Function.Magic */\n",
|
||||
"body .vc { color: #19177C } /* Name.Variable.Class */\n",
|
||||
"body .vg { color: #19177C } /* Name.Variable.Global */\n",
|
||||
"body .vi { color: #19177C } /* Name.Variable.Instance */\n",
|
||||
"body .vm { color: #19177C } /* Name.Variable.Magic */\n",
|
||||
"body .il { color: #666666 } /* Literal.Number.Integer.Long */\n",
|
||||
"\n",
|
||||
" </style>\n",
|
||||
"</head>\n",
|
||||
"<body>\n",
|
||||
"<h2></h2>\n",
|
||||
"\n",
|
||||
"<div class=\"highlight\"><pre><span></span><span class=\"k\">class</span> <span class=\"nc\">DecisionLeaf</span><span class=\"p\">:</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""A leaf of a decision tree holds just a result."""</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__init__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">result</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">result</span> <span class=\"o\">=</span> <span class=\"n\">result</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__call__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">example</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">result</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">display</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"nb\">print</span><span class=\"p\">(</span><span class=\"s1\">'RESULT ='</span><span class=\"p\">,</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">result</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__repr__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"nb\">repr</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">result</span><span class=\"p\">)</span>\n",
|
||||
"</pre></div>\n",
|
||||
"</body>\n",
|
||||
"</html>\n"
|
||||
],
|
||||
"text/plain": [
|
||||
"<IPython.core.display.HTML object>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"psource(DecisionLeaf)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "5b2935a2-15ff-4fdc-b2b3-930b1539e278",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"叶子节点在`result`中存储类别标签。所有输入图元的分类路径都在`DecisionLeaf`上结束,其`result` 属性决定了它们的类别。"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"id": "f272550b-1eac-42fd-b937-29627380b403",
|
||||
"metadata": {
|
||||
"tags": []
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/html": [
|
||||
"<!DOCTYPE html PUBLIC \"-//W3C//DTD HTML 4.01//EN\"\n",
|
||||
" \"http://www.w3.org/TR/html4/strict.dtd\">\n",
|
||||
"<!--\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"-->\n",
|
||||
"<html>\n",
|
||||
"<head>\n",
|
||||
" <title></title>\n",
|
||||
" <meta http-equiv=\"content-type\" content=\"text/html; charset=None\">\n",
|
||||
" <style type=\"text/css\">\n",
|
||||
"/*\n",
|
||||
"generated by Pygments <https://pygments.org/>\n",
|
||||
"Copyright 2006-2022 by the Pygments team.\n",
|
||||
"Licensed under the BSD license, see LICENSE for details.\n",
|
||||
"*/\n",
|
||||
"pre { line-height: 125%; }\n",
|
||||
"td.linenos .normal { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos { color: inherit; background-color: transparent; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"td.linenos .special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"span.linenos.special { color: #000000; background-color: #ffffc0; padding-left: 5px; padding-right: 5px; }\n",
|
||||
"body .hll { background-color: #ffffcc }\n",
|
||||
"body { background: #f8f8f8; }\n",
|
||||
"body .c { color: #3D7B7B; font-style: italic } /* Comment */\n",
|
||||
"body .err { border: 1px solid #FF0000 } /* Error */\n",
|
||||
"body .k { color: #008000; font-weight: bold } /* Keyword */\n",
|
||||
"body .o { color: #666666 } /* Operator */\n",
|
||||
"body .ch { color: #3D7B7B; font-style: italic } /* Comment.Hashbang */\n",
|
||||
"body .cm { color: #3D7B7B; font-style: italic } /* Comment.Multiline */\n",
|
||||
"body .cp { color: #9C6500 } /* Comment.Preproc */\n",
|
||||
"body .cpf { color: #3D7B7B; font-style: italic } /* Comment.PreprocFile */\n",
|
||||
"body .c1 { color: #3D7B7B; font-style: italic } /* Comment.Single */\n",
|
||||
"body .cs { color: #3D7B7B; font-style: italic } /* Comment.Special */\n",
|
||||
"body .gd { color: #A00000 } /* Generic.Deleted */\n",
|
||||
"body .ge { font-style: italic } /* Generic.Emph */\n",
|
||||
"body .gr { color: #E40000 } /* Generic.Error */\n",
|
||||
"body .gh { color: #000080; font-weight: bold } /* Generic.Heading */\n",
|
||||
"body .gi { color: #008400 } /* Generic.Inserted */\n",
|
||||
"body .go { color: #717171 } /* Generic.Output */\n",
|
||||
"body .gp { color: #000080; font-weight: bold } /* Generic.Prompt */\n",
|
||||
"body .gs { font-weight: bold } /* Generic.Strong */\n",
|
||||
"body .gu { color: #800080; font-weight: bold } /* Generic.Subheading */\n",
|
||||
"body .gt { color: #0044DD } /* Generic.Traceback */\n",
|
||||
"body .kc { color: #008000; font-weight: bold } /* Keyword.Constant */\n",
|
||||
"body .kd { color: #008000; font-weight: bold } /* Keyword.Declaration */\n",
|
||||
"body .kn { color: #008000; font-weight: bold } /* Keyword.Namespace */\n",
|
||||
"body .kp { color: #008000 } /* Keyword.Pseudo */\n",
|
||||
"body .kr { color: #008000; font-weight: bold } /* Keyword.Reserved */\n",
|
||||
"body .kt { color: #B00040 } /* Keyword.Type */\n",
|
||||
"body .m { color: #666666 } /* Literal.Number */\n",
|
||||
"body .s { color: #BA2121 } /* Literal.String */\n",
|
||||
"body .na { color: #687822 } /* Name.Attribute */\n",
|
||||
"body .nb { color: #008000 } /* Name.Builtin */\n",
|
||||
"body .nc { color: #0000FF; font-weight: bold } /* Name.Class */\n",
|
||||
"body .no { color: #880000 } /* Name.Constant */\n",
|
||||
"body .nd { color: #AA22FF } /* Name.Decorator */\n",
|
||||
"body .ni { color: #717171; font-weight: bold } /* Name.Entity */\n",
|
||||
"body .ne { color: #CB3F38; font-weight: bold } /* Name.Exception */\n",
|
||||
"body .nf { color: #0000FF } /* Name.Function */\n",
|
||||
"body .nl { color: #767600 } /* Name.Label */\n",
|
||||
"body .nn { color: #0000FF; font-weight: bold } /* Name.Namespace */\n",
|
||||
"body .nt { color: #008000; font-weight: bold } /* Name.Tag */\n",
|
||||
"body .nv { color: #19177C } /* Name.Variable */\n",
|
||||
"body .ow { color: #AA22FF; font-weight: bold } /* Operator.Word */\n",
|
||||
"body .w { color: #bbbbbb } /* Text.Whitespace */\n",
|
||||
"body .mb { color: #666666 } /* Literal.Number.Bin */\n",
|
||||
"body .mf { color: #666666 } /* Literal.Number.Float */\n",
|
||||
"body .mh { color: #666666 } /* Literal.Number.Hex */\n",
|
||||
"body .mi { color: #666666 } /* Literal.Number.Integer */\n",
|
||||
"body .mo { color: #666666 } /* Literal.Number.Oct */\n",
|
||||
"body .sa { color: #BA2121 } /* Literal.String.Affix */\n",
|
||||
"body .sb { color: #BA2121 } /* Literal.String.Backtick */\n",
|
||||
"body .sc { color: #BA2121 } /* Literal.String.Char */\n",
|
||||
"body .dl { color: #BA2121 } /* Literal.String.Delimiter */\n",
|
||||
"body .sd { color: #BA2121; font-style: italic } /* Literal.String.Doc */\n",
|
||||
"body .s2 { color: #BA2121 } /* Literal.String.Double */\n",
|
||||
"body .se { color: #AA5D1F; font-weight: bold } /* Literal.String.Escape */\n",
|
||||
"body .sh { color: #BA2121 } /* Literal.String.Heredoc */\n",
|
||||
"body .si { color: #A45A77; font-weight: bold } /* Literal.String.Interpol */\n",
|
||||
"body .sx { color: #008000 } /* Literal.String.Other */\n",
|
||||
"body .sr { color: #A45A77 } /* Literal.String.Regex */\n",
|
||||
"body .s1 { color: #BA2121 } /* Literal.String.Single */\n",
|
||||
"body .ss { color: #19177C } /* Literal.String.Symbol */\n",
|
||||
"body .bp { color: #008000 } /* Name.Builtin.Pseudo */\n",
|
||||
"body .fm { color: #0000FF } /* Name.Function.Magic */\n",
|
||||
"body .vc { color: #19177C } /* Name.Variable.Class */\n",
|
||||
"body .vg { color: #19177C } /* Name.Variable.Global */\n",
|
||||
"body .vi { color: #19177C } /* Name.Variable.Instance */\n",
|
||||
"body .vm { color: #19177C } /* Name.Variable.Magic */\n",
|
||||
"body .il { color: #666666 } /* Literal.Number.Integer.Long */\n",
|
||||
"\n",
|
||||
" </style>\n",
|
||||
"</head>\n",
|
||||
"<body>\n",
|
||||
"<h2></h2>\n",
|
||||
"\n",
|
||||
"<div class=\"highlight\"><pre><span></span><span class=\"k\">class</span> <span class=\"nc\">DecisionTreeLearner</span><span class=\"p\">:</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""DecisionTreeLearner: based on information gain"""</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__init__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">dataset</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span> <span class=\"o\">=</span> <span class=\"n\">dataset</span>\n",
|
||||
" <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">tree</span> <span class=\"o\">=</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">decision_tree_learning</span><span class=\"p\">(</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">examples</span><span class=\"p\">,</span> <span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">inputs</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">decision_tree_learning</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">,</span> <span class=\"n\">attrs</span><span class=\"p\">,</span> <span class=\"n\">parent_examples</span><span class=\"o\">=</span><span class=\"p\">()):</span>\n",
|
||||
" <span class=\"k\">raise</span> <span class=\"ne\">NotImplementedError</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">plurality_value</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""</span>\n",
|
||||
"<span class=\"sd\"> Return the most popular target value for this set of examples.</span>\n",
|
||||
"<span class=\"sd\"> (If target is binary, this is the majority; otherwise plurality).</span>\n",
|
||||
"<span class=\"sd\"> """</span>\n",
|
||||
" <span class=\"n\">popular</span> <span class=\"o\">=</span> <span class=\"n\">argmax_random_tie</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">values</span><span class=\"p\">[</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">target</span><span class=\"p\">],</span>\n",
|
||||
" <span class=\"n\">key</span><span class=\"o\">=</span><span class=\"k\">lambda</span> <span class=\"n\">v</span><span class=\"p\">:</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">count</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">target</span><span class=\"p\">,</span> <span class=\"n\">v</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">))</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"n\">DecisionLeaf</span><span class=\"p\">(</span><span class=\"n\">popular</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">count</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">attr</span><span class=\"p\">,</span> <span class=\"n\">val</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Count the number of examples that have example[attr] = val."""</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"nb\">sum</span><span class=\"p\">(</span><span class=\"n\">e</span><span class=\"p\">[</span><span class=\"n\">attr</span><span class=\"p\">]</span> <span class=\"o\">==</span> <span class=\"n\">val</span> <span class=\"k\">for</span> <span class=\"n\">e</span> <span class=\"ow\">in</span> <span class=\"n\">examples</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">all_same_class</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Are all these examples in the same target class?"""</span>\n",
|
||||
" <span class=\"n\">class0</span> <span class=\"o\">=</span> <span class=\"n\">examples</span><span class=\"p\">[</span><span class=\"mi\">0</span><span class=\"p\">][</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">target</span><span class=\"p\">]</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"nb\">all</span><span class=\"p\">(</span><span class=\"n\">e</span><span class=\"p\">[</span><span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">target</span><span class=\"p\">]</span> <span class=\"o\">==</span> <span class=\"n\">class0</span> <span class=\"k\">for</span> <span class=\"n\">e</span> <span class=\"ow\">in</span> <span class=\"n\">examples</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">choose_attribute</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">attrs</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Choose the attribute with the highest information gain."""</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"n\">argmax_random_tie</span><span class=\"p\">(</span><span class=\"n\">attrs</span><span class=\"p\">,</span> <span class=\"n\">key</span><span class=\"o\">=</span><span class=\"k\">lambda</span> <span class=\"n\">a</span><span class=\"p\">:</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">information_gain</span><span class=\"p\">(</span><span class=\"n\">a</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">))</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">information_gain</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">attr</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Return the expected reduction in entropy from splitting by attr."""</span>\n",
|
||||
" <span class=\"k\">raise</span> <span class=\"ne\">NotImplementedError</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">split_by</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">attr</span><span class=\"p\">,</span> <span class=\"n\">examples</span><span class=\"p\">):</span>\n",
|
||||
"<span class=\"w\"> </span><span class=\"sd\">"""Return a list of (val, examples) pairs for each val of attr."""</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"p\">[(</span><span class=\"n\">v</span><span class=\"p\">,</span> <span class=\"p\">[</span><span class=\"n\">e</span> <span class=\"k\">for</span> <span class=\"n\">e</span> <span class=\"ow\">in</span> <span class=\"n\">examples</span> <span class=\"k\">if</span> <span class=\"n\">e</span><span class=\"p\">[</span><span class=\"n\">attr</span><span class=\"p\">]</span> <span class=\"o\">==</span> <span class=\"n\">v</span><span class=\"p\">])</span> <span class=\"k\">for</span> <span class=\"n\">v</span> <span class=\"ow\">in</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">dataset</span><span class=\"o\">.</span><span class=\"n\">values</span><span class=\"p\">[</span><span class=\"n\">attr</span><span class=\"p\">]]</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"nf\">predict</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">x</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">tree</span><span class=\"p\">(</span><span class=\"n\">x</span><span class=\"p\">)</span>\n",
|
||||
"\n",
|
||||
" <span class=\"k\">def</span> <span class=\"fm\">__call__</span><span class=\"p\">(</span><span class=\"bp\">self</span><span class=\"p\">,</span> <span class=\"n\">x</span><span class=\"p\">):</span>\n",
|
||||
" <span class=\"k\">return</span> <span class=\"bp\">self</span><span class=\"o\">.</span><span class=\"n\">predict</span><span class=\"p\">(</span><span class=\"n\">x</span><span class=\"p\">)</span>\n",
|
||||
"</pre></div>\n",
|
||||
"</body>\n",
|
||||
"</html>\n"
|
||||
],
|
||||
"text/plain": [
|
||||
"<IPython.core.display.HTML object>"
|
||||
]
|
||||
},
|
||||
"metadata": {},
|
||||
"output_type": "display_data"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"psource(DecisionTreeLearner)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d22fbe3c-ed37-4c10-83da-170805f822a0",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"上面的实现使用信息增益作为衡量标准来选择测试哪一个属性进行拆分。该函数以递归的方式自上而下地构建树。根据输入,它做出四个选择中的一个。\n",
|
||||
"\n",
|
||||
"1. 如果当前步骤的输入没有训练数据,我们将返回在父步骤(上一级递归)中收到的输入数据的类别模式。\n",
|
||||
"2. 如果训练数据中的所有值都属于同一类别,它将返回一个`DecisionLeaf`,其类别标签是所有数据所属的类别。\n",
|
||||
"3. 如果数据没有可以测试的属性,我们就返回训练数据中具有最高复数值的类。\n",
|
||||
"4. 我们选择熵值最高的属性,并返回一个基于此属性的`DecisionFork`。每个分支递归地调用`decision_tree_learning`来构建子树。\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b454d72b-421e-48e8-ab0b-aa5fd81e22fb",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 实现要点\n",
|
||||
"\n",
|
||||
"```py\n",
|
||||
"def information_content(values):\n",
|
||||
" \"\"\"Number of bits to represent the probability distribution in values.\"\"\"\n",
|
||||
" probabilities = values 的归一化数值\n",
|
||||
" return probabilities 代入信息熵公式\n",
|
||||
"\n",
|
||||
"def information_gain(self, attr, examples):\n",
|
||||
" \"\"\"Return the expected reduction in entropy from splitting by attr.\"\"\"\n",
|
||||
"\n",
|
||||
" def I(examples):\n",
|
||||
" return information_content([self.count(self.dataset.target, v, examples)\n",
|
||||
" for v in self.dataset.values[self.dataset.target]])\n",
|
||||
"\n",
|
||||
" n = 样本数\n",
|
||||
" remainder = 剩余信息熵值\n",
|
||||
" return 信息增益\n",
|
||||
"```"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "f3227e10-e821-411e-b5b3-8442040b7e47",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 例子\n",
|
||||
"\n",
|
||||
"现在我们将使用决策树学习器对一个有数值的样本进行分类:5.1, 3.0, 1.1, 0.1."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"id": "d0e89769-62d8-4057-b2cc-bc165dedefac",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"setosa\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"iris = DataSet(name=\"iris\")\n",
|
||||
"DTL = DecisionTreeLearner(iris)\n",
|
||||
"print(DTL([5.1, 3.0, 1.1, 0.1]))\n",
|
||||
"#print(DTL.predict([5.1, 3.0, 1.1, 0.1]))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c923f2fe-c16c-462c-8d5d-6702094955f4",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"正如预期的那样,决策树学习器将样本归类为 \"setosa\"。"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"id": "ffc16d0d-f0c6-4f42-87b3-4f7004ce2c19",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"assert DTL.predict([5, 3, 1, 0.1]) == 'setosa'\n",
|
||||
"assert DTL.predict([6, 5, 3, 1.5]) == 'versicolor'\n",
|
||||
"assert DTL.predict([7.5, 4, 6, 2]) == 'virginica'"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.8.16"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
|
||||
|
||||
class DecisionFork:
|
||||
"""
|
||||
A fork of a decision tree holds an attribute to test, and a dict
|
||||
of branches, one for each of the attribute's values.
|
||||
"""
|
||||
|
||||
def __init__(self, attr, attr_name=None, default_child=None, branches=None):
|
||||
"""Initialize by saying what attribute this node tests."""
|
||||
self.attr = attr
|
||||
self.attr_name = attr_name or attr
|
||||
self.default_child = default_child
|
||||
self.branches = branches or {}
|
||||
|
||||
def __call__(self, example):
|
||||
"""Given an example, classify it using the attribute and the branches."""
|
||||
attr_val = example[self.attr]
|
||||
if attr_val in self.branches:
|
||||
return self.branches[attr_val](example)
|
||||
else:
|
||||
# return default class when attribute is unknown
|
||||
return self.default_child(example)
|
||||
|
||||
def add(self, val, subtree):
|
||||
"""Add a branch. If self.attr = val, go to the given subtree."""
|
||||
self.branches[val] = subtree
|
||||
|
||||
def display(self, indent=0):
|
||||
name = self.attr_name
|
||||
print('Test', name)
|
||||
for (val, subtree) in self.branches.items():
|
||||
print(' ' * 4 * indent, name, '=', val, '==>', end=' ')
|
||||
subtree.display(indent + 1)
|
||||
|
||||
def __repr__(self):
|
||||
return 'DecisionFork({0!r}, {1!r}, {2!r})'.format(self.attr, self.attr_name, self.branches)
|
||||
|
||||
|
||||
class DecisionLeaf:
|
||||
"""A leaf of a decision tree holds just a result."""
|
||||
|
||||
def __init__(self, result):
|
||||
self.result = result
|
||||
|
||||
def __call__(self, example):
|
||||
return self.result
|
||||
|
||||
def display(self):
|
||||
print('RESULT =', self.result)
|
||||
|
||||
def __repr__(self):
|
||||
return repr(self.result)
|
||||
|
||||
|
||||
class DecisionTreeLearner:
|
||||
"""DecisionTreeLearner: based on information gain"""
|
||||
|
||||
def __init__(self, dataset):
|
||||
self.dataset = dataset
|
||||
self.tree = self.decision_tree_learning(dataset.examples, dataset.inputs)
|
||||
|
||||
def decision_tree_learning(self, examples, attrs, parent_examples=()):
|
||||
if len(examples) == 0:
|
||||
return self.plurality_value(parent_examples)
|
||||
if self.all_same_class(examples):
|
||||
return DecisionLeaf(examples[0][self.dataset.target])
|
||||
if len(attrs) == 0:
|
||||
return self.plurality_value(examples)
|
||||
A = self.choose_attribute(attrs, examples)
|
||||
tree = DecisionFork(A, self.dataset.attr_names[A], self.plurality_value(examples))
|
||||
for (v_k, exs) in self.split_by(A, examples):
|
||||
subtree = self.decision_tree_learning(exs, remove_all(A, attrs), examples)
|
||||
tree.add(v_k, subtree)
|
||||
return tree
|
||||
|
||||
def plurality_value(self, examples):
|
||||
"""
|
||||
Return the most popular target value for this set of examples.
|
||||
(If target is binary, this is the majority; otherwise plurality).
|
||||
"""
|
||||
popular = argmax_random_tie(self.dataset.values[self.dataset.target],
|
||||
key=lambda v: self.count(self.dataset.target, v, examples))
|
||||
return DecisionLeaf(popular)
|
||||
|
||||
def count(self, attr, val, examples):
|
||||
"""Count the number of examples that have example[attr] = val."""
|
||||
return sum(e[attr] == val for e in examples)
|
||||
|
||||
def all_same_class(self, examples):
|
||||
"""Are all these examples in the same target class?"""
|
||||
class0 = examples[0][self.dataset.target]
|
||||
return all(e[self.dataset.target] == class0 for e in examples)
|
||||
|
||||
def choose_attribute(self, attrs, examples):
|
||||
"""Choose the attribute with the highest information gain."""
|
||||
return argmax_random_tie(attrs, key=lambda a: self.information_gain(a, examples))
|
||||
|
||||
def information_gain(self, attr, examples):
|
||||
"""Return the expected reduction in entropy from splitting by attr."""
|
||||
|
||||
def I(examples):
|
||||
return information_content([self.count(self.dataset.target, v, examples)
|
||||
for v in self.dataset.values[self.dataset.target]])
|
||||
|
||||
n = len(examples)
|
||||
remainder = sum((len(examples_i) / n) * I(examples_i)
|
||||
for (v, examples_i) in self.split_by(attr, examples))
|
||||
return I(examples) - remainder
|
||||
|
||||
def split_by(self, attr, examples):
|
||||
"""Return a list of (val, examples) pairs for each val of attr."""
|
||||
return [(v, [e for e in examples if e[attr] == v]) for v in self.dataset.values[attr]]
|
||||
|
||||
def predict(self, x):
|
||||
return self.tree(x)
|
||||
|
||||
def __call__(self, x):
|
||||
return self.predict(x)
|
||||
|
||||
|
||||
def information_content(values):
|
||||
"""Number of bits to represent the probability distribution in values."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from utils.dataset4learners import *
|
||||
|
||||
iris = DataSet(name="iris")
|
||||
DTL = DecisionTreeLearner(iris)
|
||||
print(f'DTL.predict([5, 3, 1, 0.1]): {DTL.predict([5, 3, 1, 0.1])}')
|
||||
assert DTL.predict([5, 3, 1, 0.1]) == 'setosa'
|
||||
print(f'DTL.predict([6, 5, 3, 1.5]): {DTL.predict([6, 5, 3, 1.5])}')
|
||||
assert DTL.predict([6, 5, 3, 1.5]) == 'versicolor'
|
||||
print(f'DTL.predict([7.5, 4, 6, 2]): {DTL.predict([7.5, 4, 6, 2])}')
|
||||
assert DTL.predict([7.5, 4, 6, 2]) == 'virginica'
|
||||
@@ -0,0 +1,132 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
|
||||
|
||||
class DecisionFork:
|
||||
"""
|
||||
A fork of a decision tree holds an attribute to test, and a dict
|
||||
of branches, one for each of the attribute's values.
|
||||
"""
|
||||
|
||||
def __init__(self, attr, attr_name=None, default_child=None, branches=None):
|
||||
"""Initialize by saying what attribute this node tests."""
|
||||
self.attr = attr
|
||||
self.attr_name = attr_name or attr
|
||||
self.default_child = default_child
|
||||
self.branches = branches or {}
|
||||
|
||||
def __call__(self, example):
|
||||
"""Given an example, classify it using the attribute and the branches."""
|
||||
attr_val = example[self.attr]
|
||||
if attr_val in self.branches:
|
||||
return self.branches[attr_val](example)
|
||||
else:
|
||||
# return default class when attribute is unknown
|
||||
return self.default_child(example)
|
||||
|
||||
def add(self, val, subtree):
|
||||
"""Add a branch. If self.attr = val, go to the given subtree."""
|
||||
self.branches[val] = subtree
|
||||
|
||||
def display(self, indent=0):
|
||||
name = self.attr_name
|
||||
print('Test', name)
|
||||
for (val, subtree) in self.branches.items():
|
||||
print(' ' * 4 * indent, name, '=', val, '==>', end=' ')
|
||||
subtree.display(indent + 1)
|
||||
|
||||
def __repr__(self):
|
||||
return 'DecisionFork({0!r}, {1!r}, {2!r})'.format(self.attr, self.attr_name, self.branches)
|
||||
|
||||
|
||||
class DecisionLeaf:
|
||||
"""A leaf of a decision tree holds just a result."""
|
||||
|
||||
def __init__(self, result):
|
||||
self.result = result
|
||||
|
||||
def __call__(self, example):
|
||||
return self.result
|
||||
|
||||
def display(self):
|
||||
print('RESULT =', self.result)
|
||||
|
||||
def __repr__(self):
|
||||
return repr(self.result)
|
||||
|
||||
|
||||
class DecisionTreeLearner:
|
||||
"""DecisionTreeLearner: based on information gain"""
|
||||
|
||||
def __init__(self, dataset):
|
||||
self.dataset = dataset
|
||||
self.tree = self.decision_tree_learning(dataset.examples, dataset.inputs)
|
||||
|
||||
def decision_tree_learning(self, examples, attrs, parent_examples=()):
|
||||
if len(examples) == 0:
|
||||
return self.plurality_value(parent_examples)
|
||||
if self.all_same_class(examples):
|
||||
return DecisionLeaf(examples[0][self.dataset.target])
|
||||
if len(attrs) == 0:
|
||||
return self.plurality_value(examples)
|
||||
A = self.choose_attribute(attrs, examples)
|
||||
tree = DecisionFork(A, self.dataset.attr_names[A], self.plurality_value(examples))
|
||||
for (v_k, exs) in self.split_by(A, examples):
|
||||
subtree = self.decision_tree_learning(exs, remove_all(A, attrs), examples)
|
||||
tree.add(v_k, subtree)
|
||||
return tree
|
||||
|
||||
def plurality_value(self, examples):
|
||||
"""
|
||||
Return the most popular target value for this set of examples.
|
||||
(If target is binary, this is the majority; otherwise plurality).
|
||||
"""
|
||||
popular = argmax_random_tie(self.dataset.values[self.dataset.target],
|
||||
key=lambda v: self.count(self.dataset.target, v, examples))
|
||||
return DecisionLeaf(popular)
|
||||
|
||||
def count(self, attr, val, examples):
|
||||
"""Count the number of examples that have example[attr] = val."""
|
||||
return sum(e[attr] == val for e in examples)
|
||||
|
||||
def all_same_class(self, examples):
|
||||
"""Are all these examples in the same target class?"""
|
||||
class0 = examples[0][self.dataset.target]
|
||||
return all(e[self.dataset.target] == class0 for e in examples)
|
||||
|
||||
def choose_attribute(self, attrs, examples):
|
||||
"""Choose the attribute with the highest information gain."""
|
||||
return argmax_random_tie(attrs, key=lambda a: self.information_gain(a, examples))
|
||||
|
||||
def information_gain(self, attr, examples):
|
||||
"""Return the expected reduction in entropy from splitting by attr."""
|
||||
raise NotImplementedError
|
||||
|
||||
def split_by(self, attr, examples):
|
||||
"""Return a list of (val, examples) pairs for each val of attr."""
|
||||
return [(v, [e for e in examples if e[attr] == v]) for v in self.dataset.values[attr]]
|
||||
|
||||
def predict(self, x):
|
||||
return self.tree(x)
|
||||
|
||||
def __call__(self, x):
|
||||
return self.predict(x)
|
||||
|
||||
|
||||
def information_content(values):
|
||||
"""Number of bits to represent the probability distribution in values."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from utils.dataset4learners import *
|
||||
|
||||
iris = DataSet(name="iris")
|
||||
DTL = DecisionTreeLearner(iris)
|
||||
print(f'DTL.predict([5, 3, 1, 0.1]): {DTL.predict([5, 3, 1, 0.1])}')
|
||||
assert DTL.predict([5, 3, 1, 0.1]) == 'setosa'
|
||||
print(f'DTL.predict([6, 5, 3, 1.5]): {DTL.predict([6, 5, 3, 1.5])}')
|
||||
assert DTL.predict([6, 5, 3, 1.5]) == 'versicolor'
|
||||
print(f'DTL.predict([7.5, 4, 6, 2]): {DTL.predict([7.5, 4, 6, 2])}')
|
||||
assert DTL.predict([7.5, 4, 6, 2]) == 'virginica'
|
||||
@@ -0,0 +1,121 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
|
||||
|
||||
class DecisionFork:
|
||||
"""
|
||||
A fork of a decision tree holds an attribute to test, and a dict
|
||||
of branches, one for each of the attribute's values.
|
||||
"""
|
||||
|
||||
def __init__(self, attr, attr_name=None, default_child=None, branches=None):
|
||||
"""Initialize by saying what attribute this node tests."""
|
||||
self.attr = attr
|
||||
self.attr_name = attr_name or attr
|
||||
self.default_child = default_child
|
||||
self.branches = branches or {}
|
||||
|
||||
def __call__(self, example):
|
||||
"""Given an example, classify it using the attribute and the branches."""
|
||||
attr_val = example[self.attr]
|
||||
if attr_val in self.branches:
|
||||
return self.branches[attr_val](example)
|
||||
else:
|
||||
# return default class when attribute is unknown
|
||||
return self.default_child(example)
|
||||
|
||||
def add(self, val, subtree):
|
||||
"""Add a branch. If self.attr = val, go to the given subtree."""
|
||||
self.branches[val] = subtree
|
||||
|
||||
def display(self, indent=0):
|
||||
name = self.attr_name
|
||||
print('Test', name)
|
||||
for (val, subtree) in self.branches.items():
|
||||
print(' ' * 4 * indent, name, '=', val, '==>', end=' ')
|
||||
subtree.display(indent + 1)
|
||||
|
||||
def __repr__(self):
|
||||
return 'DecisionFork({0!r}, {1!r}, {2!r})'.format(self.attr, self.attr_name, self.branches)
|
||||
|
||||
|
||||
class DecisionLeaf:
|
||||
"""A leaf of a decision tree holds just a result."""
|
||||
|
||||
def __init__(self, result):
|
||||
self.result = result
|
||||
|
||||
def __call__(self, example):
|
||||
return self.result
|
||||
|
||||
def display(self):
|
||||
print('RESULT =', self.result)
|
||||
|
||||
def __repr__(self):
|
||||
return repr(self.result)
|
||||
|
||||
|
||||
class DecisionTreeLearner:
|
||||
"""DecisionTreeLearner: based on information gain"""
|
||||
|
||||
def __init__(self, dataset):
|
||||
self.dataset = dataset
|
||||
self.tree = self.decision_tree_learning(dataset.examples, dataset.inputs)
|
||||
|
||||
def decision_tree_learning(self, examples, attrs, parent_examples=()):
|
||||
raise NotImplementedError
|
||||
|
||||
def plurality_value(self, examples):
|
||||
"""
|
||||
Return the most popular target value for this set of examples.
|
||||
(If target is binary, this is the majority; otherwise plurality).
|
||||
"""
|
||||
popular = argmax_random_tie(self.dataset.values[self.dataset.target],
|
||||
key=lambda v: self.count(self.dataset.target, v, examples))
|
||||
return DecisionLeaf(popular)
|
||||
|
||||
def count(self, attr, val, examples):
|
||||
"""Count the number of examples that have example[attr] = val."""
|
||||
return sum(e[attr] == val for e in examples)
|
||||
|
||||
def all_same_class(self, examples):
|
||||
"""Are all these examples in the same target class?"""
|
||||
class0 = examples[0][self.dataset.target]
|
||||
return all(e[self.dataset.target] == class0 for e in examples)
|
||||
|
||||
def choose_attribute(self, attrs, examples):
|
||||
"""Choose the attribute with the highest information gain."""
|
||||
return argmax_random_tie(attrs, key=lambda a: self.information_gain(a, examples))
|
||||
|
||||
def information_gain(self, attr, examples):
|
||||
"""Return the expected reduction in entropy from splitting by attr."""
|
||||
raise NotImplementedError
|
||||
|
||||
def split_by(self, attr, examples):
|
||||
"""Return a list of (val, examples) pairs for each val of attr."""
|
||||
return [(v, [e for e in examples if e[attr] == v]) for v in self.dataset.values[attr]]
|
||||
|
||||
def predict(self, x):
|
||||
return self.tree(x)
|
||||
|
||||
def __call__(self, x):
|
||||
return self.predict(x)
|
||||
|
||||
|
||||
def information_content(values):
|
||||
"""Number of bits to represent the probability distribution in values."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from utils.dataset4learners import *
|
||||
|
||||
iris = DataSet(name="iris")
|
||||
DTL = DecisionTreeLearner(iris)
|
||||
print(f'DTL.predict([5, 3, 1, 0.1]): {DTL.predict([5, 3, 1, 0.1])}')
|
||||
assert DTL.predict([5, 3, 1, 0.1]) == 'setosa'
|
||||
print(f'DTL.predict([6, 5, 3, 1.5]): {DTL.predict([6, 5, 3, 1.5])}')
|
||||
assert DTL.predict([6, 5, 3, 1.5]) == 'versicolor'
|
||||
print(f'DTL.predict([7.5, 4, 6, 2]): {DTL.predict([7.5, 4, 6, 2])}')
|
||||
assert DTL.predict([7.5, 4, 6, 2]) == 'virginica'
|
||||
@@ -0,0 +1,34 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
|
||||
|
||||
class DecisionFork:
|
||||
"""
|
||||
A fork of a decision tree holds an attribute to test, and a dict
|
||||
of branches, one for each of the attribute's values.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DecisionLeaf:
|
||||
"""A leaf of a decision tree holds just a result."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DecisionTreeLearner:
|
||||
"""DecisionTreeLearner: based on information gain"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from utils.dataset4learners import *
|
||||
|
||||
iris = DataSet(name="iris")
|
||||
DTL = DecisionTreeLearner(iris)
|
||||
print(f'DTL.predict([5, 3, 1, 0.1]): {DTL.predict([5, 3, 1, 0.1])}')
|
||||
assert DTL.predict([5, 3, 1, 0.1]) == 'setosa'
|
||||
print(f'DTL.predict([6, 5, 3, 1.5]): {DTL.predict([6, 5, 3, 1.5])}')
|
||||
assert DTL.predict([6, 5, 3, 1.5]) == 'versicolor'
|
||||
print(f'DTL.predict([7.5, 4, 6, 2]): {DTL.predict([7.5, 4, 6, 2])}')
|
||||
assert DTL.predict([7.5, 4, 6, 2]) == 'virginica'
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 45 KiB |
@@ -0,0 +1,3 @@
|
||||
DTL.predict([5, 3, 1, 0.1]): setosa
|
||||
DTL.predict([6, 5, 3, 1.5]): versicolor
|
||||
DTL.predict([7.5, 4, 6, 2]): virginica
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,147 @@
|
||||
import numpy as np
|
||||
|
||||
|
||||
class LinearRegression:
|
||||
def solve(self, lr, nepoch):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class LinearRegressionLS(LinearRegression):
|
||||
"""
|
||||
solve linear regression problem via least squares
|
||||
"""
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = self.homogeneous([x for x in range(num_data)])
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(2)
|
||||
|
||||
def homogeneous(self, xlist):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" form normal equation """
|
||||
XtX = np.dot(self.X.T, self.X)
|
||||
XtY = np.dot(self.X.T, self.Y)
|
||||
# print(XtX.shape, XtY.shape)
|
||||
self.w = np.dot(np.linalg.inv(XtX), XtY).flatten()
|
||||
# print(self.w.shape, self.w)
|
||||
|
||||
|
||||
class LinearRegressionGD1(LinearRegression):
|
||||
"""
|
||||
solve linear regression problem via gradient descent,
|
||||
using single weight vector
|
||||
"""
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = self.homogeneous([x for x in range(num_data)])
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(2)
|
||||
|
||||
def homogeneous(self, xlist):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def gd(self, lr):
|
||||
""" gradient descent update """
|
||||
def gradient(Y_hat, Y, X):
|
||||
return np.sum((Y_hat - Y) * X, axis=0)
|
||||
|
||||
Y_hat = self.linout(self.X)
|
||||
# print(Y_hat)
|
||||
grad = gradient(Y_hat, self.Y, self.X)
|
||||
self.w -= lr * grad
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" iterative solver """
|
||||
for epoch in range(num_epochs):
|
||||
self.gd(lr)
|
||||
print(f'epoch {epoch + 1}, loss {self.loss_sq(self.X, self.Y)}')
|
||||
|
||||
|
||||
class LinearRegressionGD2(LinearRegression):
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
b: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = np.array([x for x in range(num_data)]).reshape(-1, 1)
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(1)
|
||||
self.b = np.min(ylist)
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def gd(self, lr):
|
||||
""" gradient descent update """
|
||||
def gradient(Y_hat, Y, X):
|
||||
return np.array(
|
||||
[np.sum((Y_hat - Y) * X, axis=0),
|
||||
np.sum((Y_hat - Y), axis=0)
|
||||
])
|
||||
|
||||
Y_hat = self.linout(self.X)
|
||||
# print(Y_hat)
|
||||
grad = gradient(Y_hat, self.Y, self.X)
|
||||
self.w -= lr * grad[0]
|
||||
self.b -= lr * grad[1]
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" iterative solver """
|
||||
for epoch in range(num_epochs):
|
||||
self.gd(lr)
|
||||
print(f'epoch {epoch + 1}, loss {self.loss_sq(self.X, self.Y)}')
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
hp = [14213, 13448, 13870, 16192, 16415, 21501, 25910, 24866, 28981, 32926, 36741, 40974]
|
||||
hp = [x / 10000. for x in hp]
|
||||
|
||||
lr = 0.001
|
||||
num_epochs = 10
|
||||
|
||||
ls = LinearRegressionLS(hp)
|
||||
ls.solve(lr, num_epochs)
|
||||
print(f'next prediction (LS): {ls.linout([len(hp)])}')
|
||||
|
||||
gd1 = LinearRegressionGD1(hp)
|
||||
gd1.solve(lr, num_epochs)
|
||||
print(f'next prediction year (GD1): {gd1.linout([len(hp)])}')
|
||||
|
||||
gd2 = LinearRegressionGD1(hp)
|
||||
gd2.solve(lr, num_epochs)
|
||||
print(f'next prediction year (GD2): {gd2.linout([len(hp)])}')
|
||||
@@ -0,0 +1,131 @@
|
||||
import numpy as np
|
||||
|
||||
|
||||
class LinearRegression:
|
||||
def solve(self, lr, nepoch):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class LinearRegressionLS(LinearRegression):
|
||||
"""
|
||||
solve linear regression problem via least squares
|
||||
"""
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = self.homogeneous([x for x in range(num_data)])
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(2)
|
||||
|
||||
def homogeneous(self, xlist):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" form normal equation """
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class LinearRegressionGD1(LinearRegression):
|
||||
"""
|
||||
solve linear regression problem via gradient descent,
|
||||
using single weight vector: homogeneous coordinates
|
||||
"""
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = self.homogeneous([x for x in range(num_data)])
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(2)
|
||||
|
||||
def homogeneous(self, xlist):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def gd(self, lr):
|
||||
""" gradient descent update """
|
||||
raise NotImplementedError
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" iterative solver """
|
||||
for epoch in range(num_epochs):
|
||||
self.gd(lr)
|
||||
print(f'epoch {epoch + 1}, loss {self.loss_sq(self.X, self.Y)}')
|
||||
|
||||
|
||||
class LinearRegressionGD2(LinearRegression):
|
||||
"""
|
||||
solve linear regression problem via gradient descent,
|
||||
using two weights: w, b
|
||||
"""
|
||||
X: np.array
|
||||
Y: np.array
|
||||
w: np.array
|
||||
b: np.array
|
||||
|
||||
def __init__(self, ylist):
|
||||
num_data = len(ylist)
|
||||
self.X = np.array([x for x in range(num_data)]).reshape(-1, 1)
|
||||
self.Y = np.array(ylist).reshape(num_data, 1)
|
||||
self.w = np.random.rand(1)
|
||||
self.b = np.min(ylist)
|
||||
|
||||
def linout(self, xlist):
|
||||
""" linear output for given data """
|
||||
raise NotImplementedError
|
||||
|
||||
def loss_sq(self, X, Y):
|
||||
""" loss function: (half) sum of square errors """
|
||||
raise NotImplementedError
|
||||
|
||||
def gd(self, lr):
|
||||
""" gradient descent update """
|
||||
raise NotImplementedError
|
||||
|
||||
def solve(self, lr, nepoch):
|
||||
""" iterative solver """
|
||||
for epoch in range(num_epochs):
|
||||
self.gd(lr)
|
||||
print(f'epoch {epoch + 1}, loss {self.loss_sq(self.X, self.Y)}')
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
hp = [14213, 13448, 13870, 16192, 16415, 21501, 25910, 24866, 28981, 32926, 36741, 40974]
|
||||
hp = [x / 10000. for x in hp]
|
||||
|
||||
lr = 0.001
|
||||
num_epochs = 10
|
||||
|
||||
ls = LinearRegressionLS(hp)
|
||||
ls.solve(lr, num_epochs)
|
||||
print(f'next prediction (LS): {ls.linout([len(hp)])}')
|
||||
|
||||
gd1 = LinearRegressionGD1(hp)
|
||||
gd1.solve(lr, num_epochs)
|
||||
print(f'next prediction year (GD1): {gd1.linout([len(hp)])}')
|
||||
|
||||
gd2 = LinearRegressionGD1(hp)
|
||||
gd2.solve(lr, num_epochs)
|
||||
print(f'next prediction year (GD2): {gd2.linout([len(hp)])}')
|
||||
@@ -0,0 +1,38 @@
|
||||
===============================================
|
||||
next prediction (LS): [[4.04524848]]
|
||||
===============================================
|
||||
epoch 1, loss [2.25284504]
|
||||
epoch 2, loss [1.33208397]
|
||||
epoch 3, loss [1.11136605]
|
||||
epoch 4, loss [1.05556832]
|
||||
epoch 5, loss [1.03864276]
|
||||
epoch 6, loss [1.03089871]
|
||||
epoch 7, loss [1.02534236]
|
||||
epoch 8, loss [1.02032617]
|
||||
epoch 9, loss [1.01546189]
|
||||
epoch 10, loss [1.01065792]
|
||||
next prediction year (GD1): [[4.42035817]]
|
||||
===============================================
|
||||
epoch 1, loss [5.13717936]
|
||||
epoch 2, loss [1.45618137]
|
||||
epoch 3, loss [0.58898843]
|
||||
epoch 4, loss [0.38458619]
|
||||
epoch 5, loss [0.33630472]
|
||||
epoch 6, loss [0.32479828]
|
||||
epoch 7, loss [0.32195507]
|
||||
epoch 8, loss [0.32115336]
|
||||
epoch 9, loss [0.3208334]
|
||||
epoch 10, loss [0.32062779]
|
||||
next prediction year (GD2): [[4.11825378]]
|
||||
===============================================
|
||||
epoch 1, loss [0.87523181]
|
||||
epoch 2, loss [0.88078976]
|
||||
epoch 3, loss [24.9157167]
|
||||
epoch 4, loss [5055.03526329]
|
||||
epoch 5, loss [1053914.58956747]
|
||||
epoch 6, loss [2.19754618e+08]
|
||||
epoch 7, loss [4.5821661e+10]
|
||||
epoch 8, loss [9.55440499e+12]
|
||||
epoch 9, loss [1.99221619e+15]
|
||||
epoch 10, loss [4.15402668e+17]
|
||||
next prediction year (GD1): [[4.83262041e+08]]
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,129 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
from utils.dataset4learners import *
|
||||
|
||||
|
||||
class LinearClassifier:
|
||||
def learn(self, learning_rate, epochs):
|
||||
raise NotImplementedError
|
||||
|
||||
def predict(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PerceptionLinearLearner(LinearClassifier):
|
||||
"""
|
||||
Perception linear classifier: hard threshold
|
||||
"""
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, X_col, num_examples):
|
||||
""" update weights """
|
||||
for i in range(len(w)):
|
||||
w[i] = w[i] - learning_rate * (np.dot(err, X_col[i]) / num_examples)
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
X_col = homogeneous(self.num_examples)
|
||||
for epoch in range(epochs):
|
||||
err = []
|
||||
# pass over all examples
|
||||
for example in self.examples:
|
||||
err.append(loss(example, self.w, self.idx_i, self.idx_t))
|
||||
|
||||
# update weights
|
||||
update(self.w, learning_rate, err, X_col, self.num_examples)
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
class LogisticLinearLeaner(LinearClassifier):
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t, h):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, h, X_col, num_examples):
|
||||
""" update weights """
|
||||
for i in range(len(w)):
|
||||
buffer = [x * y for x, y in zip(err, h)]
|
||||
w[i] = w[i] - learning_rate * (np.dot(buffer, X_col[i]) / num_examples)
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
X_col = homogeneous(self.num_examples)
|
||||
for epoch in range(epochs):
|
||||
err = []
|
||||
h = []
|
||||
# pass over all examples
|
||||
for example in self.examples:
|
||||
err.append(loss(example, self.w, self.idx_i, self.idx_t, h))
|
||||
|
||||
# update weights
|
||||
update(self.w, learning_rate, err, h, X_col, self.num_examples)
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
iris = DataSet(name="iris")
|
||||
iris.classes_to_numbers()
|
||||
|
||||
tests = [([5, 3, 1, 0.1], 0),
|
||||
([5, 3.5, 1, 0], 0),
|
||||
([6, 3, 4, 1.1], 1),
|
||||
([6, 2, 3.5, 1], 1),
|
||||
([7.5, 4, 6, 2], 2),
|
||||
([7, 3, 6, 2.5], 2)]
|
||||
|
||||
print(f'===================\nperceptron:')
|
||||
perceptron = PerceptionLinearLearner(iris)
|
||||
g = grade_learner(perceptron, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(perceptron, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
|
||||
print(f'===================\nlogistic:')
|
||||
logisticer = LogisticLinearLeaner(iris)
|
||||
g = grade_learner(logisticer, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(logisticer, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
@@ -0,0 +1,126 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
from utils.dataset4learners import *
|
||||
|
||||
|
||||
class LinearClassifier:
|
||||
def learn(self, learning_rate, epochs):
|
||||
raise NotImplementedError
|
||||
|
||||
def predict(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PerceptionLinearLearner(LinearClassifier):
|
||||
"""
|
||||
Perception linear classifier: hard threshold
|
||||
"""
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, X_col, num_examples):
|
||||
""" update weights """
|
||||
raise NotImplementedError
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
X_col = homogeneous(self.num_examples)
|
||||
for epoch in range(epochs):
|
||||
err = []
|
||||
# pass over all examples
|
||||
for example in self.examples:
|
||||
err.append(loss(example, self.w, self.idx_i, self.idx_t))
|
||||
|
||||
# update weights
|
||||
update(self.w, learning_rate, err, X_col, self.num_examples)
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
class LogisticLinearLeaner(LinearClassifier):
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t, h):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, h, X_col, num_examples):
|
||||
""" update weights """
|
||||
raise NotImplementedError
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
X_col = homogeneous(self.num_examples)
|
||||
for epoch in range(epochs):
|
||||
err = []
|
||||
h = []
|
||||
# pass over all examples
|
||||
for example in self.examples:
|
||||
err.append(loss(example, self.w, self.idx_i, self.idx_t, h))
|
||||
|
||||
# update weights
|
||||
update(self.w, learning_rate, err, h, X_col, self.num_examples)
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
iris = DataSet(name="iris")
|
||||
iris.classes_to_numbers()
|
||||
|
||||
tests = [([5, 3, 1, 0.1], 0),
|
||||
([5, 3.5, 1, 0], 0),
|
||||
([6, 3, 4, 1.1], 1),
|
||||
([6, 2, 3.5, 1], 1),
|
||||
([7.5, 4, 6, 2], 2),
|
||||
([7, 3, 6, 2.5], 2)]
|
||||
|
||||
print(f'===================\nperceptron:')
|
||||
perceptron = PerceptionLinearLearner(iris)
|
||||
g = grade_learner(perceptron, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(perceptron, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
|
||||
print(f'===================\nlogistic:')
|
||||
logisticer = LogisticLinearLeaner(iris)
|
||||
g = grade_learner(logisticer, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(logisticer, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
@@ -0,0 +1,109 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
from utils.dataset4learners import *
|
||||
|
||||
|
||||
class LinearClassifier:
|
||||
def learn(self, learning_rate, epochs):
|
||||
raise NotImplementedError
|
||||
|
||||
def predict(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PerceptionLinearLearner(LinearClassifier):
|
||||
"""
|
||||
Perception linear classifier: hard threshold
|
||||
"""
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, X_col, num_examples):
|
||||
""" update weights """
|
||||
raise NotImplementedError
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
class LogisticLinearLeaner(LinearClassifier):
|
||||
def __init__(self, dataset, learning_rate=0.01, epochs=100):
|
||||
self.idx_i = dataset.inputs
|
||||
self.idx_t = dataset.target
|
||||
self.examples = dataset.examples
|
||||
self.num_examples = len(self.examples)
|
||||
# initialize random weights
|
||||
self.w = random_weights(min_value=-0.5, max_value=0.5, num_weights=len(self.idx_i) + 1)
|
||||
# learning loop
|
||||
self.learn(learning_rate, epochs)
|
||||
|
||||
def learn(self, learning_rate, epochs):
|
||||
""" learning loop """
|
||||
def loss(example, w, idx_i, idx_t, h):
|
||||
""" error: difference between estimation and true value """
|
||||
raise NotImplementedError
|
||||
|
||||
def update(w, learning_rate, err, h, X_col, num_examples):
|
||||
""" update weights """
|
||||
raise NotImplementedError
|
||||
|
||||
def homogeneous(num_examples):
|
||||
""" build homogeneous coordinates """
|
||||
raise NotImplementedError
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def predict(self, x):
|
||||
""" make prediction """
|
||||
return int(np.dot(self.w, [1] + x))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
iris = DataSet(name="iris")
|
||||
iris.classes_to_numbers()
|
||||
|
||||
tests = [([5, 3, 1, 0.1], 0),
|
||||
([5, 3.5, 1, 0], 0),
|
||||
([6, 3, 4, 1.1], 1),
|
||||
([6, 2, 3.5, 1], 1),
|
||||
([7.5, 4, 6, 2], 2),
|
||||
([7, 3, 6, 2.5], 2)]
|
||||
|
||||
print(f'===================\nperceptron:')
|
||||
perceptron = PerceptionLinearLearner(iris)
|
||||
g = grade_learner(perceptron, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(perceptron, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
|
||||
print(f'===================\nlogistic:')
|
||||
logisticer = LogisticLinearLeaner(iris)
|
||||
g = grade_learner(logisticer, tests)
|
||||
print(f' learner grade: {g}')
|
||||
# assert g > 1. / 2
|
||||
e = err_ratio(logisticer, iris)
|
||||
print(f' error ration: {e}')
|
||||
# assert e < 0.4
|
||||
@@ -0,0 +1,8 @@
|
||||
===================
|
||||
perceptron:
|
||||
learner grade: 0.5
|
||||
error ration: 0.3466666666666667
|
||||
===================
|
||||
logistic:
|
||||
learner grade: 0.3333333333333333
|
||||
error ration: 0.6066666666666667
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 19 KiB |
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,150 @@
|
||||
import torch
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import torch.optim as optim
|
||||
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
|
||||
|
||||
class FFNet(nn.Module):
|
||||
"""
|
||||
Feedforward neural network, virtual base class
|
||||
"""
|
||||
def __init__(self):
|
||||
super(FFNet, self).__init__()
|
||||
self.dnn_model = self.build_model()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
def train_model(self, trainloader, testloader, loss_fn, optimizer, num_epochs):
|
||||
"""Train a model."""
|
||||
|
||||
def train_epoch(model, dataloader, loss_fn, optimizer):
|
||||
"""Train a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
# Set the model to training mode
|
||||
model.train()
|
||||
for batch, (X, y) in enumerate(dataloader):
|
||||
# Compute prediction error
|
||||
pred = model(X)
|
||||
loss = loss_fn(pred, y)
|
||||
|
||||
# Backpropagation
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
if batch % 100 == 0:
|
||||
loss, current = loss.item(), batch * len(X)
|
||||
print(f"loss: {loss:>7f} [{current:>5d}/{num_data:>5d}]")
|
||||
|
||||
def test_epoch(model, dataloader, loss_fn):
|
||||
"""Test a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
num_batches = len(dataloader)
|
||||
# Set the model to evaluate mode
|
||||
model.eval()
|
||||
test_loss, correct = 0, 0
|
||||
with torch.no_grad():
|
||||
for X, y in dataloader:
|
||||
pred = model(X)
|
||||
test_loss += loss_fn(pred, y).item()
|
||||
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
|
||||
test_loss /= num_batches
|
||||
correct /= num_data
|
||||
print(f"Test Error: \n Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f} \n")
|
||||
|
||||
for epoch in range(num_epochs):
|
||||
print(f"Epoch {epoch+1}\n-------------------------------")
|
||||
train_epoch(self, trainloader, loss_fn, optimizer)
|
||||
test_epoch(self, testloader, loss_fn)
|
||||
print("Done!")
|
||||
|
||||
def make_predict(self, images):
|
||||
""" predict labels for images """
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(images)
|
||||
_, predicted = torch.max(outputs.data, 1) # (max, max_indices)
|
||||
return predicted
|
||||
|
||||
def evaluate_model(self, testloader, n):
|
||||
""" evaluation: return accuracy """
|
||||
correct = 0
|
||||
for inputs, labels in testloader:
|
||||
pred = self.make_predict(inputs)
|
||||
correct += (pred == labels).sum()
|
||||
return 100 * correct / n
|
||||
|
||||
def predict_one(self, x):
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(x)
|
||||
predicted = outputs[0].argmax(0)
|
||||
return predicted
|
||||
|
||||
|
||||
class MLP(FFNet):
|
||||
def __init__(self):
|
||||
super(MLP, self).__init__()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# load train and test set
|
||||
trainset = torchvision.datasets.MNIST(
|
||||
'../data', train=True, download=True, transform=transforms.ToTensor())
|
||||
testset = torchvision.datasets.MNIST(
|
||||
'../data', train=False, download=True, transform=transforms.ToTensor())
|
||||
print(f'Train #: {len(trainset)}; Test #: {len(testset)}')
|
||||
# data iterator
|
||||
dataiter = iter(torch.utils.data.DataLoader(trainset, batch_size=8, shuffle=False))
|
||||
images, labels = next(dataiter)
|
||||
print(f'Labels: {labels}; Batch shape: {images.size()}')
|
||||
|
||||
# construct model
|
||||
model = MLP()
|
||||
print(model)
|
||||
|
||||
# data loader: easier iteration
|
||||
BATCH_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
trainloader = torch.utils.data.DataLoader(
|
||||
trainset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
|
||||
testloader = torch.utils.data.DataLoader(
|
||||
testset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
|
||||
|
||||
# train the model
|
||||
loss_fn = nn.CrossEntropyLoss() # cross entropy
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=0.1) # SGD
|
||||
NUM_EPOCH = 5
|
||||
model.train_model(trainloader, testloader, loss_fn, optimizer, NUM_EPOCH)
|
||||
|
||||
# check prediction accuracy
|
||||
print(f'Labels : {labels}')
|
||||
print(f'Prediction: {model.make_predict(images)}')
|
||||
print(f'Accuracy: {model.evaluate_model(testloader, len(testset)):.2f}')
|
||||
|
||||
# application: predict hand-written digit
|
||||
from PIL import Image
|
||||
image = Image.open('number6c.png')
|
||||
image = transforms.ToTensor()(image).unsqueeze(0)
|
||||
print(f'loaded image shape: {image.size()}')
|
||||
print(f'Predicted: "{model.predict_one(image)}", Actual: "{6}"')
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
import torch
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import torch.optim as optim
|
||||
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
|
||||
|
||||
class FFNet(nn.Module):
|
||||
"""
|
||||
Feedforward neural network, virtual base class
|
||||
"""
|
||||
def __init__(self):
|
||||
super(FFNet, self).__init__()
|
||||
self.dnn_model = self.build_model()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
def train_model(self, trainloader, testloader, loss_fn, optimizer, num_epochs):
|
||||
"""Train a model."""
|
||||
|
||||
def train_epoch(model, dataloader, loss_fn, optimizer):
|
||||
"""Train a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
# Set the model to training mode
|
||||
model.train()
|
||||
for batch, (X, y) in enumerate(dataloader):
|
||||
# Compute prediction error
|
||||
pred = model(X)
|
||||
loss = loss_fn(pred, y)
|
||||
|
||||
# Backpropagation
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
if batch % 100 == 0:
|
||||
loss, current = loss.item(), batch * len(X)
|
||||
print(f"loss: {loss:>7f} [{current:>5d}/{num_data:>5d}]")
|
||||
|
||||
def test_epoch(model, dataloader, loss_fn):
|
||||
"""Test a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
num_batches = len(dataloader)
|
||||
# Set the model to evaluate mode
|
||||
model.eval()
|
||||
test_loss, correct = 0, 0
|
||||
with torch.no_grad():
|
||||
for X, y in dataloader:
|
||||
pred = model(X)
|
||||
test_loss += loss_fn(pred, y).item()
|
||||
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
|
||||
test_loss /= num_batches
|
||||
correct /= num_data
|
||||
print(f"Test Error: \n Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f} \n")
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def make_predict(self, images):
|
||||
""" predict labels for images """
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(images)
|
||||
_, predicted = torch.max(outputs.data, 1) # (max, max_indices)
|
||||
return predicted
|
||||
|
||||
def evaluate_model(self, testloader, n):
|
||||
""" evaluation: return accuracy """
|
||||
raise NotImplementedError
|
||||
|
||||
def predict_one(self, x):
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(x)
|
||||
predicted = outputs[0].argmax(0)
|
||||
return predicted
|
||||
|
||||
|
||||
class MLP(FFNet):
|
||||
def __init__(self):
|
||||
super(MLP, self).__init__()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# load train and test set
|
||||
trainset = torchvision.datasets.MNIST(
|
||||
'../data', train=True, download=True, transform=transforms.ToTensor())
|
||||
testset = torchvision.datasets.MNIST(
|
||||
'../data', train=False, download=True, transform=transforms.ToTensor())
|
||||
print(f'Train #: {len(trainset)}; Test #: {len(testset)}')
|
||||
# data iterator
|
||||
dataiter = iter(torch.utils.data.DataLoader(trainset, batch_size=8, shuffle=False))
|
||||
images, labels = next(dataiter)
|
||||
print(f'Labels: {labels}; Batch shape: {images.size()}')
|
||||
|
||||
# construct model
|
||||
model = MLP()
|
||||
print(model)
|
||||
|
||||
# data loader: easier iteration
|
||||
BATCH_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
trainloader = torch.utils.data.DataLoader(
|
||||
trainset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
|
||||
testloader = torch.utils.data.DataLoader(
|
||||
testset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
|
||||
|
||||
# train the model
|
||||
loss_fn = nn.CrossEntropyLoss() # cross entropy
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=0.1) # SGD
|
||||
NUM_EPOCH = 5
|
||||
model.train_model(trainloader, testloader, loss_fn, optimizer, NUM_EPOCH)
|
||||
|
||||
# check prediction accuracy
|
||||
print(f'Labels : {labels}')
|
||||
print(f'Prediction: {model.make_predict(images)}')
|
||||
print(f'Accuracy: {model.evaluate_model(testloader, len(testset)):.2f}')
|
||||
|
||||
# application: predict hand-written digit
|
||||
from PIL import Image
|
||||
image = Image.open('number6c.png')
|
||||
image = transforms.ToTensor()(image).unsqueeze(0)
|
||||
print(f'loaded image shape: {image.size()}')
|
||||
print(f'Predicted: "{model.predict_one(image)}", Actual: "{6}"')
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
import torch
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import torch.optim as optim
|
||||
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
|
||||
|
||||
class FFNet(nn.Module):
|
||||
"""
|
||||
Feedforward neural network, virtual base class
|
||||
"""
|
||||
def __init__(self):
|
||||
super(FFNet, self).__init__()
|
||||
self.dnn_model = self.build_model()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
def train_model(self, trainloader, testloader, loss_fn, optimizer, num_epochs):
|
||||
"""Train a model."""
|
||||
|
||||
def train_epoch(model, dataloader, loss_fn, optimizer):
|
||||
"""Train a single epoch"""
|
||||
raise NotImplementedError
|
||||
|
||||
def test_epoch(model, dataloader, loss_fn):
|
||||
"""Test a single epoch"""
|
||||
raise NotImplementedError
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def make_predict(self, images):
|
||||
""" predict labels for images """
|
||||
raise NotImplementedError
|
||||
|
||||
def evaluate_model(self, testloader, n):
|
||||
""" evaluation: return accuracy """
|
||||
raise NotImplementedError
|
||||
|
||||
def predict_one(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MLP(FFNet):
|
||||
def __init__(self):
|
||||
super(MLP, self).__init__()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# load train and test set
|
||||
trainset = torchvision.datasets.MNIST(
|
||||
'../data', train=True, download=True, transform=transforms.ToTensor())
|
||||
testset = torchvision.datasets.MNIST(
|
||||
'../data', train=False, download=True, transform=transforms.ToTensor())
|
||||
print(f'Train #: {len(trainset)}; Test #: {len(testset)}')
|
||||
# data iterator
|
||||
dataiter = iter(torch.utils.data.DataLoader(trainset, batch_size=8, shuffle=False))
|
||||
images, labels = next(dataiter)
|
||||
print(f'Labels: {labels}; Batch shape: {images.size()}')
|
||||
|
||||
# construct model
|
||||
model = MLP()
|
||||
print(model)
|
||||
|
||||
# data loader: easier iteration
|
||||
BATCH_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
trainloader = torch.utils.data.DataLoader(
|
||||
trainset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
|
||||
testloader = torch.utils.data.DataLoader(
|
||||
testset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
|
||||
|
||||
# train the model
|
||||
loss_fn = nn.CrossEntropyLoss() # cross entropy
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=0.1) # SGD
|
||||
NUM_EPOCH = 5
|
||||
model.train_model(trainloader, testloader, loss_fn, optimizer, NUM_EPOCH)
|
||||
|
||||
# check prediction accuracy
|
||||
print(f'Labels : {labels}')
|
||||
print(f'Prediction: {model.make_predict(images)}')
|
||||
print(f'Accuracy: {model.evaluate_model(testloader, len(testset)):.2f}')
|
||||
|
||||
# application: predict hand-written digit
|
||||
from PIL import Image
|
||||
image = Image.open('number6c.png')
|
||||
image = transforms.ToTensor()(image).unsqueeze(0)
|
||||
print(f'loaded image shape: {image.size()}')
|
||||
print(f'Predicted: "{model.predict_one(image)}", Actual: "{6}"')
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
Train #: 60000; Test #: 10000
|
||||
Labels: tensor([5, 0, 4, 1, 9, 2, 1, 3]); Batch shape: torch.Size([8, 1, 28, 28])
|
||||
MLP(
|
||||
(flatten): Flatten(start_dim=1, end_dim=-1)
|
||||
(dnn_model): Sequential(
|
||||
(0): Linear(in_features=784, out_features=256, bias=True)
|
||||
(1): ReLU()
|
||||
(2): Linear(in_features=256, out_features=10, bias=True)
|
||||
)
|
||||
)
|
||||
Epoch 1
|
||||
-------------------------------
|
||||
loss: 2.308738 [ 0/60000]
|
||||
loss: 0.643267 [25600/60000]
|
||||
loss: 0.384817 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 90.0%, Avg loss: 0.368403
|
||||
|
||||
Epoch 2
|
||||
-------------------------------
|
||||
loss: 0.323617 [ 0/60000]
|
||||
loss: 0.300336 [25600/60000]
|
||||
loss: 0.268706 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 91.6%, Avg loss: 0.301624
|
||||
|
||||
Epoch 3
|
||||
-------------------------------
|
||||
loss: 0.314411 [ 0/60000]
|
||||
loss: 0.337723 [25600/60000]
|
||||
loss: 0.234745 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 92.4%, Avg loss: 0.267629
|
||||
|
||||
Epoch 4
|
||||
-------------------------------
|
||||
loss: 0.258774 [ 0/60000]
|
||||
loss: 0.222343 [25600/60000]
|
||||
loss: 0.363658 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 92.9%, Avg loss: 0.245184
|
||||
|
||||
Epoch 5
|
||||
-------------------------------
|
||||
loss: 0.365407 [ 0/60000]
|
||||
loss: 0.286917 [25600/60000]
|
||||
loss: 0.213272 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 93.9%, Avg loss: 0.222085
|
||||
|
||||
Done!
|
||||
Labels : tensor([5, 0, 4, 1, 9, 2, 1, 3])
|
||||
Prediction: tensor([5, 0, 4, 1, 9, 2, 1, 3])
|
||||
Accuracy: 93.87
|
||||
loaded image shape: torch.Size([1, 1, 28, 28])
|
||||
Predicted: "3", Actual: "6"
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 176 B |
Binary file not shown.
|
After Width: | Height: | Size: 169 B |
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,150 @@
|
||||
import torch
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import torch.optim as optim
|
||||
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
|
||||
|
||||
class FFNet(nn.Module):
|
||||
"""
|
||||
Feedforward neural network, virtual base class
|
||||
"""
|
||||
def __init__(self):
|
||||
super(FFNet, self).__init__()
|
||||
self.dnn_model = self.build_model()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
raise NotImplementedError
|
||||
|
||||
def train_model(self, trainloader, testloader, loss_fn, optimizer, num_epochs):
|
||||
"""Train a model."""
|
||||
|
||||
def train_epoch(model, dataloader, loss_fn, optimizer):
|
||||
"""Train a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
# Set the model to training mode
|
||||
model.train()
|
||||
for batch, (X, y) in enumerate(dataloader):
|
||||
# Compute prediction error
|
||||
pred = model(X)
|
||||
loss = loss_fn(pred, y)
|
||||
|
||||
# Backpropagation
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
if batch % 100 == 0:
|
||||
loss, current = loss.item(), batch * len(X)
|
||||
print(f"loss: {loss:>7f} [{current:>5d}/{num_data:>5d}]")
|
||||
|
||||
def test_epoch(model, dataloader, loss_fn):
|
||||
"""Test a single epoch"""
|
||||
num_data = len(dataloader.dataset)
|
||||
num_batches = len(dataloader)
|
||||
# Set the model to evaluate mode
|
||||
model.eval()
|
||||
test_loss, correct = 0, 0
|
||||
with torch.no_grad():
|
||||
for X, y in dataloader:
|
||||
pred = model(X)
|
||||
test_loss += loss_fn(pred, y).item()
|
||||
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
|
||||
test_loss /= num_batches
|
||||
correct /= num_data
|
||||
print(f"Test Error: \n Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f} \n")
|
||||
|
||||
for epoch in range(num_epochs):
|
||||
print(f"Epoch {epoch+1}\n-------------------------------")
|
||||
train_epoch(self, trainloader, loss_fn, optimizer)
|
||||
test_epoch(self, testloader, loss_fn)
|
||||
print("Done!")
|
||||
|
||||
def make_predict(self, images):
|
||||
""" predict labels for images """
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(images)
|
||||
_, predicted = torch.max(outputs.data, 1) # (max, max_indices)
|
||||
return predicted
|
||||
|
||||
def evaluate_model(self, testloader, n):
|
||||
""" evaluation: return accuracy """
|
||||
correct = 0
|
||||
for inputs, labels in testloader:
|
||||
pred = self.make_predict(inputs)
|
||||
correct += (pred == labels).sum()
|
||||
return 100 * correct / n
|
||||
|
||||
def predict_one(self, x):
|
||||
self.eval()
|
||||
with torch.no_grad():
|
||||
outputs = self(x)
|
||||
predicted = outputs[0].argmax(0)
|
||||
return predicted
|
||||
|
||||
|
||||
class LeNet(FFNet):
|
||||
def __init__(self):
|
||||
super(LeNet, self).__init__()
|
||||
|
||||
def build_model(self):
|
||||
""" build specific model """
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, x):
|
||||
""" feed data forward and return result """
|
||||
return self.dnn_model(x)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# load train and test set
|
||||
trainset = torchvision.datasets.MNIST(
|
||||
'../data', train=True, download=True, transform=transforms.ToTensor())
|
||||
testset = torchvision.datasets.MNIST(
|
||||
'../data', train=False, download=True, transform=transforms.ToTensor())
|
||||
print(f'Train #: {len(trainset)}; Test #: {len(testset)}')
|
||||
# data iterator
|
||||
dataiter = iter(torch.utils.data.DataLoader(trainset, batch_size=8, shuffle=False))
|
||||
images, labels = next(dataiter)
|
||||
print(f'Labels: {labels}; Batch shape: {images.size()}')
|
||||
|
||||
# construct model
|
||||
model = LeNet()
|
||||
print(model)
|
||||
|
||||
# data loader: easier iteration
|
||||
BATCH_SIZE = 256
|
||||
NUM_WORKERS = 4
|
||||
trainloader = torch.utils.data.DataLoader(
|
||||
trainset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS)
|
||||
testloader = torch.utils.data.DataLoader(
|
||||
testset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)
|
||||
|
||||
# train the model
|
||||
LR = 0.1
|
||||
loss_fn = nn.CrossEntropyLoss() # cross entropy
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=LR) # SGD
|
||||
NUM_EPOCH = 5
|
||||
model.train_model(trainloader, testloader, loss_fn, optimizer, NUM_EPOCH)
|
||||
|
||||
# check prediction accuracy
|
||||
print(f'Labels : {labels}')
|
||||
print(f'Prediction: {model.make_predict(images)}')
|
||||
print(f'Accuracy: {model.evaluate_model(testloader, len(testset)):.2f}')
|
||||
|
||||
# application: predict hand-written digit
|
||||
from PIL import Image
|
||||
image = Image.open('number6c.png')
|
||||
image = transforms.ToTensor()(image).unsqueeze(0)
|
||||
print(f'loaded image shape: {image.size()}')
|
||||
print(f'Predicted: "{model.predict_one(image)}", Actual: "{6}"')
|
||||
@@ -0,0 +1,64 @@
|
||||
Train #: 60000; Test #: 10000
|
||||
Labels: tensor([5, 0, 4, 1, 9, 2, 1, 3]); Batch shape: torch.Size([8, 1, 28, 28])
|
||||
LeNet(
|
||||
(dnn_model): Sequential(
|
||||
(0): Conv2d(1, 32, kernel_size=(5, 5), stride=(1, 1))
|
||||
(1): ReLU()
|
||||
(2): Conv2d(32, 32, kernel_size=(5, 5), stride=(1, 1))
|
||||
(3): ReLU()
|
||||
(4): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
|
||||
(5): Conv2d(32, 64, kernel_size=(5, 5), stride=(1, 1))
|
||||
(6): ReLU()
|
||||
(7): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
|
||||
(8): Flatten(start_dim=1, end_dim=-1)
|
||||
(9): Linear(in_features=576, out_features=256, bias=True)
|
||||
(10): ReLU()
|
||||
(11): Linear(in_features=256, out_features=10, bias=True)
|
||||
)
|
||||
)
|
||||
Epoch 1
|
||||
-------------------------------
|
||||
loss: 2.299683 [ 0/60000]
|
||||
loss: 0.348185 [25600/60000]
|
||||
loss: 0.239826 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 95.7%, Avg loss: 0.133152
|
||||
|
||||
Epoch 2
|
||||
-------------------------------
|
||||
loss: 0.220472 [ 0/60000]
|
||||
loss: 0.056995 [25600/60000]
|
||||
loss: 0.078447 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 90.3%, Avg loss: 0.284392
|
||||
|
||||
Epoch 3
|
||||
-------------------------------
|
||||
loss: 0.255085 [ 0/60000]
|
||||
loss: 0.071809 [25600/60000]
|
||||
loss: 0.080565 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 98.2%, Avg loss: 0.056539
|
||||
|
||||
Epoch 4
|
||||
-------------------------------
|
||||
loss: 0.074963 [ 0/60000]
|
||||
loss: 0.053749 [25600/60000]
|
||||
loss: 0.042180 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 98.3%, Avg loss: 0.048151
|
||||
|
||||
Epoch 5
|
||||
-------------------------------
|
||||
loss: 0.018707 [ 0/60000]
|
||||
loss: 0.047416 [25600/60000]
|
||||
loss: 0.022978 [51200/60000]
|
||||
Test Error:
|
||||
Accuracy: 98.6%, Avg loss: 0.041758
|
||||
|
||||
Done!
|
||||
Labels : tensor([5, 0, 4, 1, 9, 2, 1, 3])
|
||||
Prediction: tensor([5, 0, 4, 1, 9, 2, 1, 3])
|
||||
Accuracy: 98.60
|
||||
loaded image shape: torch.Size([1, 1, 28, 28])
|
||||
Predicted: "0", Actual: "6"
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 176 B |
Binary file not shown.
|
After Width: | Height: | Size: 169 B |
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,218 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
|
||||
class RNN(nn.Module):
|
||||
""" Recurrent neural network - simple
|
||||
x -> hidden -> out -> log-softmax
|
||||
"""
|
||||
def __init__(self, input_size, hidden_size, output_size):
|
||||
super(RNN, self).__init__()
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, input, hidden):
|
||||
raise NotImplementedError
|
||||
|
||||
def initHidden(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
##########################################################
|
||||
# Text data utilities - character-level
|
||||
import glob
|
||||
import unicodedata
|
||||
import string
|
||||
|
||||
|
||||
all_letters = string.ascii_letters + " .,;'-" # all legal letters
|
||||
n_letters = len(all_letters)
|
||||
|
||||
def unicodeToAscii(s):
|
||||
""" Unicode string to plain ASCII,
|
||||
ref. http://stackoverflow.com/a/518232/2809427
|
||||
"""
|
||||
return ''.join(
|
||||
c for c in unicodedata.normalize('NFD', s)
|
||||
if unicodedata.category(c) != 'Mn'
|
||||
and c in all_letters
|
||||
)
|
||||
|
||||
def readLines(filename):
|
||||
""" Read a file and split into lines """
|
||||
lines = open(filename).read().strip().split('\n')
|
||||
return [unicodeToAscii(line) for line in lines]
|
||||
|
||||
def build_category_lines(data_file='../data/names/*.txt'):
|
||||
""" Build the category_lines dictionary, a list of lines per category """
|
||||
|
||||
def findFiles(path): return glob.glob(path)
|
||||
|
||||
category_lines = {}
|
||||
all_categories = []
|
||||
for filename in findFiles(data_file):
|
||||
category = filename.split('/')[-1].split('.')[0]
|
||||
all_categories.append(category)
|
||||
lines = readLines(filename)
|
||||
category_lines[category] = lines
|
||||
|
||||
return category_lines, all_categories
|
||||
|
||||
def letterToIndex(letter):
|
||||
""" Find letter index from all_letters, e.g. "a" = 0 """
|
||||
return all_letters.find(letter)
|
||||
|
||||
def lineToTensor(line):
|
||||
""" Turn a line into a <line_length x 1 x n_letters>,
|
||||
ie, an array of one-hot letter vectors
|
||||
"""
|
||||
tensor = torch.zeros(len(line), 1, n_letters)
|
||||
for li, letter in enumerate(line):
|
||||
tensor[li][0][letterToIndex(letter)] = 1
|
||||
return tensor
|
||||
##########################################################
|
||||
|
||||
|
||||
import random
|
||||
import time
|
||||
import math
|
||||
|
||||
|
||||
class TextLearner:
|
||||
@staticmethod
|
||||
def train_epoch(model, category_tensor, line_tensor):
|
||||
hidden = model.initHidden()
|
||||
optimizer.zero_grad()
|
||||
|
||||
for i in range(line_tensor.size()[0]):
|
||||
output, hidden = model(line_tensor[i], hidden)
|
||||
|
||||
loss = criterion(output, category_tensor)
|
||||
loss.backward()
|
||||
|
||||
optimizer.step()
|
||||
|
||||
return output, loss.item()
|
||||
|
||||
@staticmethod
|
||||
def train(model, n_epochs, print_every, plot_every, learning_rate):
|
||||
def categoryFromOutput(output):
|
||||
top_n, top_i = output.data.topk(1) # Tensor out of Variable with .data
|
||||
category_i = top_i[0][0]
|
||||
return all_categories[category_i], category_i
|
||||
|
||||
def randomTrainingPair():
|
||||
""" generate a random training pair """
|
||||
def randomChoice(l): return l[random.randint(0, len(l) - 1)]
|
||||
|
||||
category = randomChoice(all_categories)
|
||||
line = randomChoice(category_lines[category])
|
||||
category_tensor = Variable(torch.LongTensor([all_categories.index(category)]))
|
||||
line_tensor = Variable(lineToTensor(line))
|
||||
return category, line, category_tensor, line_tensor
|
||||
|
||||
def timeSince(since):
|
||||
now = time.time()
|
||||
s = now - since
|
||||
m = math.floor(s / 60)
|
||||
s -= m * 60
|
||||
return f'{m}m {int(s): >2d}s'
|
||||
start = time.time()
|
||||
|
||||
# Keep track of losses for plotting
|
||||
current_loss = 0
|
||||
all_losses = []
|
||||
for epoch in range(1, n_epochs + 1):
|
||||
category, line, category_tensor, line_tensor = randomTrainingPair()
|
||||
output, loss = TextLearner.train_epoch(model, category_tensor, line_tensor)
|
||||
current_loss += loss
|
||||
|
||||
# Print epoch number, loss, name and guess
|
||||
if epoch % print_every == 0:
|
||||
guess, guess_i = categoryFromOutput(output)
|
||||
correct = '✓' if guess == category else '✗ (%s)' % category
|
||||
print(f'{epoch: >6d} {epoch / n_epochs * 100:5.1f}% ({timeSince(start)}) {loss:.4f} {current_loss:.2f} {line} / {guess} {correct}')
|
||||
|
||||
# Add current loss avg to list of losses
|
||||
if epoch % plot_every == 0:
|
||||
all_losses.append(current_loss / plot_every)
|
||||
current_loss = 0
|
||||
|
||||
torch.save(model, 'char_rnn_names.pt')
|
||||
|
||||
@staticmethod
|
||||
def predict_t(line_tensor):
|
||||
""" return an output given a line """
|
||||
hidden = model.initHidden()
|
||||
|
||||
for i in range(line_tensor.size()[0]):
|
||||
output, hidden = model(line_tensor[i], hidden)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def predict(model, line, n_predictions=3):
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
|
||||
# Get top N categories
|
||||
topv, topi = output.data.topk(n_predictions, 1, True)
|
||||
predictions = []
|
||||
|
||||
print(f'prediction for {line}:')
|
||||
for i in range(n_predictions):
|
||||
value = topv[0][i]
|
||||
category_index = topi[0][i]
|
||||
print(f' ({value:.2f}) {all_categories[category_index]}')
|
||||
predictions.append([value, all_categories[category_index]])
|
||||
|
||||
return predictions
|
||||
|
||||
@staticmethod
|
||||
def eval_all(model, n_predictions=5):
|
||||
correct_n = []
|
||||
total_n = []
|
||||
correct_str = ''
|
||||
for ci, category in enumerate(category_lines):
|
||||
total_n.append(len(category))
|
||||
cni = 0
|
||||
for line in category:
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
_, topi = output.data.topk(n_predictions)
|
||||
cc = 0
|
||||
for ii in range(n_predictions):
|
||||
if (topi[0][ii] == ci):
|
||||
cc = 1
|
||||
break
|
||||
cni += cc
|
||||
correct_n.append(cni)
|
||||
correct_str += f'{all_categories[ci]}: {cni / total_n[-1] * 100:5.1f}; '
|
||||
|
||||
print(f'correct rate (per category): {correct_str}')
|
||||
print(f'total correct rate: {sum(correct_n) / sum(total_n) * 100:5.1f}')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
print(f'{n_letters} legal letters: {all_letters}')
|
||||
category_lines, all_categories = build_category_lines()
|
||||
n_categories = len(all_categories)
|
||||
print(f'{n_categories} categories: {all_categories}')
|
||||
print(f"first 5 Chinese names: {category_lines['Chinese'][:5]}")
|
||||
|
||||
n_hidden = 128
|
||||
n_epochs = 100000
|
||||
print_every = 5000
|
||||
plot_every = 1000
|
||||
learning_rate = 0.005
|
||||
|
||||
model = RNN(n_letters, n_hidden, n_categories)
|
||||
print(model)
|
||||
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate)
|
||||
criterion = nn.NLLLoss() # negative log likelihood loss
|
||||
TextLearner.train(model, n_epochs, print_every, plot_every, learning_rate)
|
||||
|
||||
model = torch.load('char_rnn_names.pt')
|
||||
TextLearner.predict(model, 'Wu')
|
||||
TextLearner.predict(model, 'Harry')
|
||||
TextLearner.predict(model, 'Louis')
|
||||
|
||||
TextLearner.eval_all(model)
|
||||
@@ -0,0 +1,215 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
|
||||
class RNN(nn.Module):
|
||||
""" Recurrent neural network - simple
|
||||
x -> hidden -> out -> log-softmax
|
||||
"""
|
||||
def __init__(self, input_size, hidden_size, output_size):
|
||||
super(RNN, self).__init__()
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, input, hidden):
|
||||
raise NotImplementedError
|
||||
|
||||
def initHidden(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
##########################################################
|
||||
# Text data utilities - character-level
|
||||
import glob
|
||||
import unicodedata
|
||||
import string
|
||||
|
||||
|
||||
all_letters = string.ascii_letters + " .,;'-" # all legal letters
|
||||
n_letters = len(all_letters)
|
||||
|
||||
def unicodeToAscii(s):
|
||||
""" Unicode string to plain ASCII,
|
||||
ref. http://stackoverflow.com/a/518232/2809427
|
||||
"""
|
||||
return ''.join(
|
||||
c for c in unicodedata.normalize('NFD', s)
|
||||
if unicodedata.category(c) != 'Mn'
|
||||
and c in all_letters
|
||||
)
|
||||
|
||||
def readLines(filename):
|
||||
""" Read a file and split into lines """
|
||||
lines = open(filename).read().strip().split('\n')
|
||||
return [unicodeToAscii(line) for line in lines]
|
||||
|
||||
def build_category_lines(data_file='../data/names/*.txt'):
|
||||
""" Build the category_lines dictionary, a list of lines per category """
|
||||
|
||||
def findFiles(path): return glob.glob(path)
|
||||
|
||||
category_lines = {}
|
||||
all_categories = []
|
||||
for filename in findFiles(data_file):
|
||||
category = filename.split('/')[-1].split('.')[0]
|
||||
all_categories.append(category)
|
||||
lines = readLines(filename)
|
||||
category_lines[category] = lines
|
||||
|
||||
return category_lines, all_categories
|
||||
|
||||
def letterToIndex(letter):
|
||||
""" Find letter index from all_letters, e.g. "a" = 0 """
|
||||
raise NotImplementedError
|
||||
|
||||
def lineToTensor(line):
|
||||
""" Turn a line into a <line_length x 1 x n_letters>,
|
||||
ie, an array of one-hot letter vectors
|
||||
"""
|
||||
raise NotImplementedError
|
||||
##########################################################
|
||||
|
||||
|
||||
import random
|
||||
import time
|
||||
import math
|
||||
|
||||
|
||||
class TextLearner:
|
||||
@staticmethod
|
||||
def train_epoch(model, category_tensor, line_tensor):
|
||||
hidden = model.initHidden()
|
||||
optimizer.zero_grad()
|
||||
|
||||
for i in range(line_tensor.size()[0]):
|
||||
output, hidden = model(line_tensor[i], hidden)
|
||||
|
||||
loss = criterion(output, category_tensor)
|
||||
loss.backward()
|
||||
|
||||
optimizer.step()
|
||||
|
||||
return output, loss.item()
|
||||
|
||||
@staticmethod
|
||||
def train(model, n_epochs, print_every, plot_every, learning_rate):
|
||||
def categoryFromOutput(output):
|
||||
top_n, top_i = output.data.topk(1) # Tensor out of Variable with .data
|
||||
category_i = top_i[0][0]
|
||||
return all_categories[category_i], category_i
|
||||
|
||||
def randomTrainingPair():
|
||||
""" generate a random training pair """
|
||||
def randomChoice(l): return l[random.randint(0, len(l) - 1)]
|
||||
|
||||
category = randomChoice(all_categories)
|
||||
line = randomChoice(category_lines[category])
|
||||
category_tensor = Variable(torch.LongTensor([all_categories.index(category)]))
|
||||
line_tensor = Variable(lineToTensor(line))
|
||||
return category, line, category_tensor, line_tensor
|
||||
|
||||
def timeSince(since):
|
||||
now = time.time()
|
||||
s = now - since
|
||||
m = math.floor(s / 60)
|
||||
s -= m * 60
|
||||
return f'{m}m {int(s): >2d}s'
|
||||
start = time.time()
|
||||
|
||||
# Keep track of losses for plotting
|
||||
current_loss = 0
|
||||
all_losses = []
|
||||
for epoch in range(1, n_epochs + 1):
|
||||
category, line, category_tensor, line_tensor = randomTrainingPair()
|
||||
output, loss = TextLearner.train_epoch(model, category_tensor, line_tensor)
|
||||
current_loss += loss
|
||||
|
||||
# Print epoch number, loss, name and guess
|
||||
if epoch % print_every == 0:
|
||||
guess, guess_i = categoryFromOutput(output)
|
||||
correct = '✓' if guess == category else '✗ (%s)' % category
|
||||
print(f'{epoch: >6d} {epoch / n_epochs * 100:5.1f}% ({timeSince(start)}) {loss:.4f} {current_loss:.2f} {line} / {guess} {correct}')
|
||||
|
||||
# Add current loss avg to list of losses
|
||||
if epoch % plot_every == 0:
|
||||
all_losses.append(current_loss / plot_every)
|
||||
current_loss = 0
|
||||
|
||||
torch.save(model, 'char_rnn_names.pt')
|
||||
|
||||
@staticmethod
|
||||
def predict_t(line_tensor):
|
||||
""" return an output given a line """
|
||||
hidden = model.initHidden()
|
||||
|
||||
for i in range(line_tensor.size()[0]):
|
||||
output, hidden = model(line_tensor[i], hidden)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def predict(model, line, n_predictions=3):
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
|
||||
# Get top N categories
|
||||
topv, topi = output.data.topk(n_predictions, 1, True)
|
||||
predictions = []
|
||||
|
||||
print(f'prediction for {line}:')
|
||||
for i in range(n_predictions):
|
||||
value = topv[0][i]
|
||||
category_index = topi[0][i]
|
||||
print(f' ({value:.2f}) {all_categories[category_index]}')
|
||||
predictions.append([value, all_categories[category_index]])
|
||||
|
||||
return predictions
|
||||
|
||||
@staticmethod
|
||||
def eval_all(model, n_predictions=5):
|
||||
correct_n = []
|
||||
total_n = []
|
||||
correct_str = ''
|
||||
for ci, category in enumerate(category_lines):
|
||||
total_n.append(len(category))
|
||||
cni = 0
|
||||
for line in category:
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
_, topi = output.data.topk(n_predictions)
|
||||
cc = 0
|
||||
for ii in range(n_predictions):
|
||||
if (topi[0][ii] == ci):
|
||||
cc = 1
|
||||
break
|
||||
cni += cc
|
||||
correct_n.append(cni)
|
||||
correct_str += f'{all_categories[ci]}: {cni / total_n[-1] * 100:5.1f}; '
|
||||
|
||||
print(f'correct rate (per category): {correct_str}')
|
||||
print(f'total correct rate: {sum(correct_n) / sum(total_n) * 100:5.1f}')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
print(f'{n_letters} legal letters: {all_letters}')
|
||||
category_lines, all_categories = build_category_lines()
|
||||
n_categories = len(all_categories)
|
||||
print(f'{n_categories} categories: {all_categories}')
|
||||
print(f"first 5 Chinese names: {category_lines['Chinese'][:5]}")
|
||||
|
||||
n_hidden = 128
|
||||
n_epochs = 100000
|
||||
print_every = 5000
|
||||
plot_every = 1000
|
||||
learning_rate = 0.005
|
||||
|
||||
model = RNN(n_letters, n_hidden, n_categories)
|
||||
print(model)
|
||||
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate)
|
||||
criterion = nn.NLLLoss() # negative log likelihood loss
|
||||
TextLearner.train(model, n_epochs, print_every, plot_every, learning_rate)
|
||||
|
||||
model = torch.load('char_rnn_names.pt')
|
||||
TextLearner.predict(model, 'Wu')
|
||||
TextLearner.predict(model, 'Harry')
|
||||
TextLearner.predict(model, 'Louis')
|
||||
|
||||
TextLearner.eval_all(model)
|
||||
@@ -0,0 +1,199 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Variable
|
||||
|
||||
class RNN(nn.Module):
|
||||
""" Recurrent neural network - simple
|
||||
x -> hidden -> out -> log-softmax
|
||||
"""
|
||||
def __init__(self, input_size, hidden_size, output_size):
|
||||
super(RNN, self).__init__()
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, input, hidden):
|
||||
raise NotImplementedError
|
||||
|
||||
def initHidden(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
##########################################################
|
||||
# Text data utilities - character-level
|
||||
import glob
|
||||
import unicodedata
|
||||
import string
|
||||
|
||||
|
||||
all_letters = string.ascii_letters + " .,;'-" # all legal letters
|
||||
n_letters = len(all_letters)
|
||||
|
||||
def unicodeToAscii(s):
|
||||
""" Unicode string to plain ASCII,
|
||||
ref. http://stackoverflow.com/a/518232/2809427
|
||||
"""
|
||||
return ''.join(
|
||||
c for c in unicodedata.normalize('NFD', s)
|
||||
if unicodedata.category(c) != 'Mn'
|
||||
and c in all_letters
|
||||
)
|
||||
|
||||
def readLines(filename):
|
||||
""" Read a file and split into lines """
|
||||
lines = open(filename).read().strip().split('\n')
|
||||
return [unicodeToAscii(line) for line in lines]
|
||||
|
||||
def build_category_lines(data_file='../data/names/*.txt'):
|
||||
""" Build the category_lines dictionary, a list of lines per category """
|
||||
|
||||
def findFiles(path): return glob.glob(path)
|
||||
|
||||
category_lines = {}
|
||||
all_categories = []
|
||||
for filename in findFiles(data_file):
|
||||
category = filename.split('/')[-1].split('.')[0]
|
||||
all_categories.append(category)
|
||||
lines = readLines(filename)
|
||||
category_lines[category] = lines
|
||||
|
||||
return category_lines, all_categories
|
||||
|
||||
def letterToIndex(letter):
|
||||
""" Find letter index from all_letters, e.g. "a" = 0 """
|
||||
raise NotImplementedError
|
||||
|
||||
def lineToTensor(line):
|
||||
""" Turn a line into a <line_length x 1 x n_letters>,
|
||||
ie, an array of one-hot letter vectors
|
||||
"""
|
||||
raise NotImplementedError
|
||||
##########################################################
|
||||
|
||||
|
||||
import random
|
||||
import time
|
||||
import math
|
||||
|
||||
|
||||
class TextLearner:
|
||||
@staticmethod
|
||||
def train_epoch(model, category_tensor, line_tensor):
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def train(model, n_epochs, print_every, plot_every, learning_rate):
|
||||
def categoryFromOutput(output):
|
||||
top_n, top_i = output.data.topk(1) # Tensor out of Variable with .data
|
||||
category_i = top_i[0][0]
|
||||
return all_categories[category_i], category_i
|
||||
|
||||
def randomTrainingPair():
|
||||
""" generate a random training pair """
|
||||
def randomChoice(l): return l[random.randint(0, len(l) - 1)]
|
||||
|
||||
category = randomChoice(all_categories)
|
||||
line = randomChoice(category_lines[category])
|
||||
category_tensor = Variable(torch.LongTensor([all_categories.index(category)]))
|
||||
line_tensor = Variable(lineToTensor(line))
|
||||
return category, line, category_tensor, line_tensor
|
||||
|
||||
def timeSince(since):
|
||||
now = time.time()
|
||||
s = now - since
|
||||
m = math.floor(s / 60)
|
||||
s -= m * 60
|
||||
return f'{m}m {int(s): >2d}s'
|
||||
start = time.time()
|
||||
|
||||
# Keep track of losses for plotting
|
||||
current_loss = 0
|
||||
all_losses = []
|
||||
for epoch in range(1, n_epochs + 1):
|
||||
category, line, category_tensor, line_tensor = randomTrainingPair()
|
||||
output, loss = TextLearner.train_epoch(model, category_tensor, line_tensor)
|
||||
current_loss += loss
|
||||
|
||||
# Print epoch number, loss, name and guess
|
||||
if epoch % print_every == 0:
|
||||
guess, guess_i = categoryFromOutput(output)
|
||||
correct = '✓' if guess == category else '✗ (%s)' % category
|
||||
print(f'{epoch: >6d} {epoch / n_epochs * 100:5.1f}% ({timeSince(start)}) {loss:.4f} {current_loss:.2f} {line} / {guess} {correct}')
|
||||
|
||||
# Add current loss avg to list of losses
|
||||
if epoch % plot_every == 0:
|
||||
all_losses.append(current_loss / plot_every)
|
||||
current_loss = 0
|
||||
|
||||
torch.save(model, 'char_rnn_names.pt')
|
||||
|
||||
@staticmethod
|
||||
def predict_t(line_tensor):
|
||||
""" return an output given a line """
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def predict(model, line, n_predictions=3):
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
|
||||
# Get top N categories
|
||||
topv, topi = output.data.topk(n_predictions, 1, True)
|
||||
predictions = []
|
||||
|
||||
print(f'prediction for {line}:')
|
||||
for i in range(n_predictions):
|
||||
value = topv[0][i]
|
||||
category_index = topi[0][i]
|
||||
print(f' ({value:.2f}) {all_categories[category_index]}')
|
||||
predictions.append([value, all_categories[category_index]])
|
||||
|
||||
return predictions
|
||||
|
||||
@staticmethod
|
||||
def eval_all(model, n_predictions=5):
|
||||
correct_n = []
|
||||
total_n = []
|
||||
correct_str = ''
|
||||
for ci, category in enumerate(category_lines):
|
||||
total_n.append(len(category))
|
||||
cni = 0
|
||||
for line in category:
|
||||
output = TextLearner.predict_t(Variable(lineToTensor(line)))
|
||||
_, topi = output.data.topk(n_predictions)
|
||||
cc = 0
|
||||
for ii in range(n_predictions):
|
||||
if (topi[0][ii] == ci):
|
||||
cc = 1
|
||||
break
|
||||
cni += cc
|
||||
correct_n.append(cni)
|
||||
correct_str += f'{all_categories[ci]}: {cni / total_n[-1] * 100:5.1f}; '
|
||||
|
||||
print(f'correct rate (per category): {correct_str}')
|
||||
print(f'total correct rate: {sum(correct_n) / sum(total_n) * 100:5.1f}')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
print(f'{n_letters} legal letters: {all_letters}')
|
||||
category_lines, all_categories = build_category_lines()
|
||||
n_categories = len(all_categories)
|
||||
print(f'{n_categories} categories: {all_categories}')
|
||||
print(f"first 5 Chinese names: {category_lines['Chinese'][:5]}")
|
||||
|
||||
n_hidden = 128
|
||||
n_epochs = 100000
|
||||
print_every = 5000
|
||||
plot_every = 1000
|
||||
learning_rate = 0.005
|
||||
|
||||
model = RNN(n_letters, n_hidden, n_categories)
|
||||
print(model)
|
||||
|
||||
optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate)
|
||||
criterion = nn.NLLLoss() # negative log likelihood loss
|
||||
TextLearner.train(model, n_epochs, print_every, plot_every, learning_rate)
|
||||
|
||||
model = torch.load('char_rnn_names.pt')
|
||||
TextLearner.predict(model, 'Wu')
|
||||
TextLearner.predict(model, 'Harry')
|
||||
TextLearner.predict(model, 'Louis')
|
||||
|
||||
TextLearner.eval_all(model)
|
||||
Binary file not shown.
@@ -0,0 +1,42 @@
|
||||
58 legal letters: abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ .,;'-
|
||||
18 categories: ['Czech', 'German', 'Arabic', 'Japanese', 'Chinese', 'Vietnamese', 'Russian', 'French', 'Irish', 'English', 'Spanish', 'Greek', 'Italian', 'Portuguese', 'Scottish', 'Dutch', 'Korean', 'Polish']
|
||||
first 5 Chinese names: ['Ang', 'Au-Yong', 'Bai', 'Ban', 'Bao']
|
||||
RNN(
|
||||
(i2h): Linear(in_features=186, out_features=128, bias=True)
|
||||
(i2o): Linear(in_features=186, out_features=18, bias=True)
|
||||
(softmax): LogSoftmax(dim=1)
|
||||
)
|
||||
5000 5.0% (0m 9s) 2.5912 2541.35 Zimmerman / Dutch ✗ (German)
|
||||
10000 10.0% (0m 18s) 1.5788 2128.60 Gong / Vietnamese ✗ (Chinese)
|
||||
15000 15.0% (0m 26s) 2.7233 1998.65 Solo / Chinese ✗ (Spanish)
|
||||
20000 20.0% (0m 36s) 1.8938 1815.65 Lohrenz / Spanish ✗ (German)
|
||||
25000 25.0% (0m 44s) 2.3676 1735.06 Hoch / Vietnamese ✗ (German)
|
||||
30000 30.0% (0m 52s) 1.8657 1634.36 Macfarland / French ✗ (Irish)
|
||||
35000 35.0% (1m 4s) 1.5567 1658.71 Devin / French ✗ (Irish)
|
||||
40000 40.0% (1m 14s) 3.7105 1568.38 Farmer / French ✗ (English)
|
||||
45000 45.0% (1m 22s) 1.9464 1524.05 Arthur / Arabic ✗ (French)
|
||||
50000 50.0% (1m 32s) 2.1889 1463.57 Nemec / Portuguese ✗ (Czech)
|
||||
55000 55.0% (1m 42s) 0.6579 1500.78 Naser / Arabic ✓
|
||||
60000 60.0% (1m 50s) 0.7514 1427.26 Gallego / Spanish ✓
|
||||
65000 65.0% (1m 58s) 0.8220 1383.95 Kawate / Japanese ✓
|
||||
70000 70.0% (2m 6s) 0.7611 1319.17 Graner / German ✓
|
||||
75000 75.0% (2m 14s) 1.2482 1331.15 Guerrero / Spanish ✓
|
||||
80000 80.0% (2m 23s) 0.0109 1306.75 Antoniadis / Greek ✓
|
||||
85000 85.0% (2m 31s) 0.0735 1404.86 O'Boyle / Irish ✓
|
||||
90000 90.0% (2m 39s) 0.8547 1258.60 Said / Arabic ✓
|
||||
95000 95.0% (2m 47s) 1.2185 1329.16 Dickson / English ✗ (Scottish)
|
||||
100000 100.0% (2m 55s) 1.6481 1295.43 Abana / Italian ✗ (Spanish)
|
||||
prediction for Wu:
|
||||
(-0.70) Korean
|
||||
(-1.36) Vietnamese
|
||||
(-2.53) Chinese
|
||||
prediction for Harry:
|
||||
(-1.27) Arabic
|
||||
(-1.63) English
|
||||
(-1.81) French
|
||||
prediction for Louis:
|
||||
(-1.16) Arabic
|
||||
(-1.45) Greek
|
||||
(-1.58) Portuguese
|
||||
correct rate (per category): Czech: 40.0; German: 66.7; Arabic: 33.3; Japanese: 25.0; Chinese: 14.3; Vietnamese: 30.0; Russian: 0.0; French: 50.0; Irish: 20.0; English: 85.7; Spanish: 14.3; Greek: 0.0; Italian: 14.3; Portuguese: 20.0; Scottish: 25.0; Dutch: 0.0; Korean: 66.7; Polish: 0.0;
|
||||
total correct rate: 28.1
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,5 @@
|
||||
NumRooms,Alley,Price
|
||||
NA,Pave,127500
|
||||
2,NA,106000
|
||||
4,NA,178100
|
||||
NA,NA,140000
|
||||
|
+150
@@ -0,0 +1,150 @@
|
||||
5.1,3.5,1.4,0.2,setosa
|
||||
4.9,3.0,1.4,0.2,setosa
|
||||
4.7,3.2,1.3,0.2,setosa
|
||||
4.6,3.1,1.5,0.2,setosa
|
||||
5.0,3.6,1.4,0.2,setosa
|
||||
5.4,3.9,1.7,0.4,setosa
|
||||
4.6,3.4,1.4,0.3,setosa
|
||||
5.0,3.4,1.5,0.2,setosa
|
||||
4.4,2.9,1.4,0.2,setosa
|
||||
4.9,3.1,1.5,0.1,setosa
|
||||
5.4,3.7,1.5,0.2,setosa
|
||||
4.8,3.4,1.6,0.2,setosa
|
||||
4.8,3.0,1.4,0.1,setosa
|
||||
4.3,3.0,1.1,0.1,setosa
|
||||
5.8,4.0,1.2,0.2,setosa
|
||||
5.7,4.4,1.5,0.4,setosa
|
||||
5.4,3.9,1.3,0.4,setosa
|
||||
5.1,3.5,1.4,0.3,setosa
|
||||
5.7,3.8,1.7,0.3,setosa
|
||||
5.1,3.8,1.5,0.3,setosa
|
||||
5.4,3.4,1.7,0.2,setosa
|
||||
5.1,3.7,1.5,0.4,setosa
|
||||
4.6,3.6,1.0,0.2,setosa
|
||||
5.1,3.3,1.7,0.5,setosa
|
||||
4.8,3.4,1.9,0.2,setosa
|
||||
5.0,3.0,1.6,0.2,setosa
|
||||
5.0,3.4,1.6,0.4,setosa
|
||||
5.2,3.5,1.5,0.2,setosa
|
||||
5.2,3.4,1.4,0.2,setosa
|
||||
4.7,3.2,1.6,0.2,setosa
|
||||
4.8,3.1,1.6,0.2,setosa
|
||||
5.4,3.4,1.5,0.4,setosa
|
||||
5.2,4.1,1.5,0.1,setosa
|
||||
5.5,4.2,1.4,0.2,setosa
|
||||
4.9,3.1,1.5,0.1,setosa
|
||||
5.0,3.2,1.2,0.2,setosa
|
||||
5.5,3.5,1.3,0.2,setosa
|
||||
4.9,3.1,1.5,0.1,setosa
|
||||
4.4,3.0,1.3,0.2,setosa
|
||||
5.1,3.4,1.5,0.2,setosa
|
||||
5.0,3.5,1.3,0.3,setosa
|
||||
4.5,2.3,1.3,0.3,setosa
|
||||
4.4,3.2,1.3,0.2,setosa
|
||||
5.0,3.5,1.6,0.6,setosa
|
||||
5.1,3.8,1.9,0.4,setosa
|
||||
4.8,3.0,1.4,0.3,setosa
|
||||
5.1,3.8,1.6,0.2,setosa
|
||||
4.6,3.2,1.4,0.2,setosa
|
||||
5.3,3.7,1.5,0.2,setosa
|
||||
5.0,3.3,1.4,0.2,setosa
|
||||
7.0,3.2,4.7,1.4,versicolor
|
||||
6.4,3.2,4.5,1.5,versicolor
|
||||
6.9,3.1,4.9,1.5,versicolor
|
||||
5.5,2.3,4.0,1.3,versicolor
|
||||
6.5,2.8,4.6,1.5,versicolor
|
||||
5.7,2.8,4.5,1.3,versicolor
|
||||
6.3,3.3,4.7,1.6,versicolor
|
||||
4.9,2.4,3.3,1.0,versicolor
|
||||
6.6,2.9,4.6,1.3,versicolor
|
||||
5.2,2.7,3.9,1.4,versicolor
|
||||
5.0,2.0,3.5,1.0,versicolor
|
||||
5.9,3.0,4.2,1.5,versicolor
|
||||
6.0,2.2,4.0,1.0,versicolor
|
||||
6.1,2.9,4.7,1.4,versicolor
|
||||
5.6,2.9,3.6,1.3,versicolor
|
||||
6.7,3.1,4.4,1.4,versicolor
|
||||
5.6,3.0,4.5,1.5,versicolor
|
||||
5.8,2.7,4.1,1.0,versicolor
|
||||
6.2,2.2,4.5,1.5,versicolor
|
||||
5.6,2.5,3.9,1.1,versicolor
|
||||
5.9,3.2,4.8,1.8,versicolor
|
||||
6.1,2.8,4.0,1.3,versicolor
|
||||
6.3,2.5,4.9,1.5,versicolor
|
||||
6.1,2.8,4.7,1.2,versicolor
|
||||
6.4,2.9,4.3,1.3,versicolor
|
||||
6.6,3.0,4.4,1.4,versicolor
|
||||
6.8,2.8,4.8,1.4,versicolor
|
||||
6.7,3.0,5.0,1.7,versicolor
|
||||
6.0,2.9,4.5,1.5,versicolor
|
||||
5.7,2.6,3.5,1.0,versicolor
|
||||
5.5,2.4,3.8,1.1,versicolor
|
||||
5.5,2.4,3.7,1.0,versicolor
|
||||
5.8,2.7,3.9,1.2,versicolor
|
||||
6.0,2.7,5.1,1.6,versicolor
|
||||
5.4,3.0,4.5,1.5,versicolor
|
||||
6.0,3.4,4.5,1.6,versicolor
|
||||
6.7,3.1,4.7,1.5,versicolor
|
||||
6.3,2.3,4.4,1.3,versicolor
|
||||
5.6,3.0,4.1,1.3,versicolor
|
||||
5.5,2.5,4.0,1.3,versicolor
|
||||
5.5,2.6,4.4,1.2,versicolor
|
||||
6.1,3.0,4.6,1.4,versicolor
|
||||
5.8,2.6,4.0,1.2,versicolor
|
||||
5.0,2.3,3.3,1.0,versicolor
|
||||
5.6,2.7,4.2,1.3,versicolor
|
||||
5.7,3.0,4.2,1.2,versicolor
|
||||
5.7,2.9,4.2,1.3,versicolor
|
||||
6.2,2.9,4.3,1.3,versicolor
|
||||
5.1,2.5,3.0,1.1,versicolor
|
||||
5.7,2.8,4.1,1.3,versicolor
|
||||
6.3,3.3,6.0,2.5,virginica
|
||||
5.8,2.7,5.1,1.9,virginica
|
||||
7.1,3.0,5.9,2.1,virginica
|
||||
6.3,2.9,5.6,1.8,virginica
|
||||
6.5,3.0,5.8,2.2,virginica
|
||||
7.6,3.0,6.6,2.1,virginica
|
||||
4.9,2.5,4.5,1.7,virginica
|
||||
7.3,2.9,6.3,1.8,virginica
|
||||
6.7,2.5,5.8,1.8,virginica
|
||||
7.2,3.6,6.1,2.5,virginica
|
||||
6.5,3.2,5.1,2.0,virginica
|
||||
6.4,2.7,5.3,1.9,virginica
|
||||
6.8,3.0,5.5,2.1,virginica
|
||||
5.7,2.5,5.0,2.0,virginica
|
||||
5.8,2.8,5.1,2.4,virginica
|
||||
6.4,3.2,5.3,2.3,virginica
|
||||
6.5,3.0,5.5,1.8,virginica
|
||||
7.7,3.8,6.7,2.2,virginica
|
||||
7.7,2.6,6.9,2.3,virginica
|
||||
6.0,2.2,5.0,1.5,virginica
|
||||
6.9,3.2,5.7,2.3,virginica
|
||||
5.6,2.8,4.9,2.0,virginica
|
||||
7.7,2.8,6.7,2.0,virginica
|
||||
6.3,2.7,4.9,1.8,virginica
|
||||
6.7,3.3,5.7,2.1,virginica
|
||||
7.2,3.2,6.0,1.8,virginica
|
||||
6.2,2.8,4.8,1.8,virginica
|
||||
6.1,3.0,4.9,1.8,virginica
|
||||
6.4,2.8,5.6,2.1,virginica
|
||||
7.2,3.0,5.8,1.6,virginica
|
||||
7.4,2.8,6.1,1.9,virginica
|
||||
7.9,3.8,6.4,2.0,virginica
|
||||
6.4,2.8,5.6,2.2,virginica
|
||||
6.3,2.8,5.1,1.5,virginica
|
||||
6.1,2.6,5.6,1.4,virginica
|
||||
7.7,3.0,6.1,2.3,virginica
|
||||
6.3,3.4,5.6,2.4,virginica
|
||||
6.4,3.1,5.5,1.8,virginica
|
||||
6.0,3.0,4.8,1.8,virginica
|
||||
6.9,3.1,5.4,2.1,virginica
|
||||
6.7,3.1,5.6,2.4,virginica
|
||||
6.9,3.1,5.1,2.3,virginica
|
||||
5.8,2.7,5.1,1.9,virginica
|
||||
6.8,3.2,5.9,2.3,virginica
|
||||
6.7,3.3,5.7,2.5,virginica
|
||||
6.7,3.0,5.2,2.3,virginica
|
||||
6.3,2.5,5.0,1.9,virginica
|
||||
6.5,3.0,5.2,2.0,virginica
|
||||
6.2,3.4,5.4,2.3,virginica
|
||||
5.9,3.0,5.1,1.8,virginica
|
||||
|
@@ -0,0 +1,69 @@
|
||||
1. Title: Iris Plants Database
|
||||
Updated Sept 21 by C.Blake - Added discrepency information
|
||||
|
||||
2. Sources:
|
||||
(a) Creator: R.A. Fisher
|
||||
(b) Donor: Michael Marshall (MARSHALL%PLU@io.arc.nasa.gov)
|
||||
(c) Date: July, 1988
|
||||
|
||||
3. Past Usage:
|
||||
- Publications: too many to mention!!! Here are a few.
|
||||
1. Fisher,R.A. "The use of multiple measurements in taxonomic problems"
|
||||
Annual Eugenics, 7, Part II, 179-188 (1936); also in "Contributions
|
||||
to Mathematical Statistics" (John Wiley, NY, 1950).
|
||||
2. Duda,R.O., & Hart,P.E. (1973) Pattern Classification and Scene Analysis.
|
||||
(Q327.D83) John Wiley & Sons. ISBN 0-471-22361-1. See page 218.
|
||||
3. Dasarathy, B.V. (1980) "Nosing Around the Neighborhood: A New System
|
||||
Structure and Classification Rule for Recognition in Partially Exposed
|
||||
Environments". IEEE Transactions on Pattern Analysis and Machine
|
||||
Intelligence, Vol. PAMI-2, No. 1, 67-71.
|
||||
-- Results:
|
||||
-- very low misclassification rates (0% for the setosa class)
|
||||
4. Gates, G.W. (1972) "The Reduced Nearest Neighbor Rule". IEEE
|
||||
Transactions on Information Theory, May 1972, 431-433.
|
||||
-- Results:
|
||||
-- very low misclassification rates again
|
||||
5. See also: 1988 MLC Proceedings, 54-64. Cheeseman et al's AUTOCLASS II
|
||||
conceptual clustering system finds 3 classes in the data.
|
||||
|
||||
4. Relevant Information:
|
||||
--- This is perhaps the best known database to be found in the pattern
|
||||
recognition literature. Fisher's paper is a classic in the field
|
||||
and is referenced frequently to this day. (See Duda & Hart, for
|
||||
example.) The data set contains 3 classes of 50 instances each,
|
||||
where each class refers to a type of iris plant. One class is
|
||||
linearly separable from the other 2; the latter are NOT linearly
|
||||
separable from each other.
|
||||
--- Predicted attribute: class of iris plant.
|
||||
--- This is an exceedingly simple domain.
|
||||
--- This data differs from the data presented in Fishers article
|
||||
(identified by Steve Chadwick, spchadwick@espeedaz.net )
|
||||
The 35th sample should be: 4.9,3.1,1.5,0.2,"Iris-setosa"
|
||||
where the error is in the fourth feature.
|
||||
The 38th sample: 4.9,3.6,1.4,0.1,"Iris-setosa"
|
||||
where the errors are in the second and third features.
|
||||
|
||||
5. Number of Instances: 150 (50 in each of three classes)
|
||||
|
||||
6. Number of Attributes: 4 numeric, predictive attributes and the class
|
||||
|
||||
7. Attribute Information:
|
||||
1. sepal length in cm
|
||||
2. sepal width in cm
|
||||
3. petal length in cm
|
||||
4. petal width in cm
|
||||
5. class:
|
||||
-- Iris Setosa
|
||||
-- Iris Versicolour
|
||||
-- Iris Virginica
|
||||
|
||||
8. Missing Attribute Values: None
|
||||
|
||||
Summary Statistics:
|
||||
Min Max Mean SD Class Correlation
|
||||
sepal length: 4.3 7.9 5.84 0.83 0.7826
|
||||
sepal width: 2.0 4.4 3.05 0.43 -0.4194
|
||||
petal length: 1.0 6.9 3.76 1.76 0.9490 (high!)
|
||||
petal width: 0.1 2.5 1.20 0.76 0.9565 (high!)
|
||||
|
||||
9. Class Distribution: 33.3% for each of 3 classes.
|
||||
@@ -0,0 +1,23 @@
|
||||
6, 0, 66, 50, 1
|
||||
6, 1, 70, 50, 2
|
||||
6, 0, 69, 50, 3
|
||||
6, 0, 68, 50, 4
|
||||
6, 0, 67, 50, 5
|
||||
6, 0, 72, 50, 6
|
||||
6, 0, 73, 100, 7
|
||||
6, 0, 70, 100, 8
|
||||
6, 1, 57, 200, 9
|
||||
6, 1, 63, 200, 10
|
||||
6, 1, 70, 200, 11
|
||||
6, 0, 78, 200, 12
|
||||
6, 0, 67, 200, 13
|
||||
6, 2, 53, 200, 14
|
||||
6, 0, 67, 200, 15
|
||||
6, 0, 75, 200, 16
|
||||
6, 0, 70, 200, 17
|
||||
6, 0, 81, 200, 18
|
||||
6, 0, 76, 200, 19
|
||||
6, 0, 79, 200, 20
|
||||
6, 0, 75, 200, 21
|
||||
6, 0, 76, 200, 22
|
||||
6, 1, 58, 200, 23
|
||||
|
@@ -0,0 +1,78 @@
|
||||
Source: http://www1.ics.uci.edu/pub/machine-learning-databases/space-shuttle/
|
||||
|
||||
1. Title: Challenger Space Shuttle O-Ring Data (2 databases)
|
||||
|
||||
2. Sources:
|
||||
-- David Draper (draper@math.ucla.edu)
|
||||
University of California, Los Angeles
|
||||
-- Donor: David Draper (draper@math.ucla.edu)
|
||||
-- Date: 5 August 1993
|
||||
|
||||
3. Past Usage:
|
||||
|
||||
1. Draper,~D. (1993). Assessment and propagation of model uncertainty.
|
||||
In {\it Proceedings of the Fourth International Workshop on Artificial
|
||||
Intelligence and Statistics} (pp. 497--509). Ft. Lauderdale, FL:
|
||||
Unpublished.
|
||||
-- Discrete model uncertainty analysis
|
||||
-- Analysis suggests that obvious different extrapolations of the
|
||||
data exist at 31 degrees Fahrenheit (i.e., freezing), which sharply
|
||||
discredits the assumption of no temperature effect.
|
||||
2. Dalal,~S.~R., Fowlkes,~E.~B., \& Hoadley,~B. (1989). Risk analysis of
|
||||
the space shuttle: pre-Challenger prediction of failure. {\it Journal
|
||||
of the American Statisticians Association}, {\it 84}, 945--957.
|
||||
3. Lavine,~M. (1991). Problems in extrapolation illustrated with space
|
||||
shuttle O-ring data. {\it Journal of the American Statisticians
|
||||
Association}, {\it 86}, 919--922.
|
||||
4. Martz~H.~F., \& Zimmer,~W.~J. (1992). The risk of catastrophic failure
|
||||
of the solid rocket boosters on the space shuttle. {\it American
|
||||
Statistics}, {\it 46}, 42--47.
|
||||
|
||||
4. Number of instances: 23 in each of two files
|
||||
|
||||
5. Relevant Information:
|
||||
|
||||
There are two databases: (both use the same set of 5 attributes)
|
||||
1. Primary o-ring erosion and/or blowby
|
||||
2. Primary o-ring erosion only
|
||||
The two databases are identical except for the 2nd attribute of the
|
||||
21st instance (confirmed by David Draper on 8/5/93).
|
||||
|
||||
Edited from (Draper, 1993):
|
||||
The motivation for collecting this database was the explosion of the
|
||||
USA Space Shuttle Challenger on 28 January, 1986. An investigation
|
||||
ensued into the reliability of the shuttle's propulsion system. The
|
||||
explosion was eventually traced to the failure of one of the three field
|
||||
joints on one of the two solid booster rockets. Each of these six field
|
||||
joints includes two O-rings, designated as primary and secondary, which
|
||||
fail when phenomena called erosion and blowby both occur.
|
||||
The night before the launch a decision had to be made regarding
|
||||
launch safety. The discussion among engineers and managers leading to
|
||||
this decision included concern that the probability of failure of the
|
||||
O-rings depended on the temperature t at launch, which was forecase to
|
||||
be 31 degrees F. There are strong engineering reasons based on the
|
||||
composition of O-rings to support the judgment that failure
|
||||
probability may rise monotonically as temperature drops. One other
|
||||
variable, the pressure s at which safety testing for field join leaks
|
||||
was performed, was available, but its relevance to the failure process
|
||||
was unclear.
|
||||
Draper's paper includes a menacing figure graphing the number of field
|
||||
joints experiencing stress vs. liftoff temperature for the 23 shuttle
|
||||
flights previous to the Challenger disaster. No previous liftoff
|
||||
temperature was under 53 degrees F. Although tremendous extrapolation
|
||||
must be done from the given data to assess risk at 31 degrees F, it
|
||||
is obvious even to the layman "to foresee the unacceptably high risk
|
||||
created by launching at 31 degrees F." For more information, see
|
||||
Draper (1993) or the other previous analyses.
|
||||
The task is to predict the number of O-rings that will experience
|
||||
thermal distress for a given flight when the launch temperature is
|
||||
below freezing.
|
||||
|
||||
6. Number of Attributes: 5
|
||||
1. Number of O-rings at risk on a given flight
|
||||
2. Number experiencing thermal distress
|
||||
3. Launch temperature (degrees F)
|
||||
4. Leak-check pressure (psi)
|
||||
5. Temporal order of flight
|
||||
|
||||
7. Attribute Information: all values are positive integers
|
||||
@@ -0,0 +1,12 @@
|
||||
Yes, No, No, Yes, Some, $$$, No, Yes, French, 0-10, Yes
|
||||
Yes, No, No, Yes, Full, $, No, No, Thai, 30-60, No
|
||||
No, Yes, No, No, Some, $, No, No, Burger, 0-10, Yes
|
||||
Yes, No, Yes, Yes, Full, $, No, No, Thai, 10-30, Yes
|
||||
Yes, No, Yes, No, Full, $$$, No, Yes, French, >60, No
|
||||
No, Yes, No, Yes, Some, $$, Yes, Yes, Italian, 0-10, Yes
|
||||
No, Yes, No, No, None, $, Yes, No, Burger, 0-10, No
|
||||
No, No, No, Yes, Some, $$, Yes, Yes, Thai, 0-10, Yes
|
||||
No, Yes, Yes, No, Full, $, Yes, No, Burger, >60, No
|
||||
Yes, Yes, Yes, Yes, Full, $$$, No, Yes, Italian, 10-30, No
|
||||
No, No, No, No, None, $, No, No, Thai, 0-10, No
|
||||
Yes, Yes, Yes, Yes, Full, $, No, No, Burger, 30-60, Yes
|
||||
|
+101
@@ -0,0 +1,101 @@
|
||||
aardvark,1,0,0,1,0,0,1,1,1,1,0,0,4,0,0,1,mammal
|
||||
antelope,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
bass,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,0,fish
|
||||
bear,1,0,0,1,0,0,1,1,1,1,0,0,4,0,0,1,mammal
|
||||
boar,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
buffalo,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
calf,1,0,0,1,0,0,0,1,1,1,0,0,4,1,1,1,mammal
|
||||
carp,0,0,1,0,0,1,0,1,1,0,0,1,0,1,1,0,fish
|
||||
catfish,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,0,fish
|
||||
cavy,1,0,0,1,0,0,0,1,1,1,0,0,4,0,1,0,mammal
|
||||
cheetah,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
chicken,0,1,1,0,1,0,0,0,1,1,0,0,2,1,1,0,bird
|
||||
chub,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,0,fish
|
||||
clam,0,0,1,0,0,0,1,0,0,0,0,0,0,0,0,0,shellfish
|
||||
crab,0,0,1,0,0,1,1,0,0,0,0,0,4,0,0,0,shellfish
|
||||
crayfish,0,0,1,0,0,1,1,0,0,0,0,0,6,0,0,0,shellfish
|
||||
crow,0,1,1,0,1,0,1,0,1,1,0,0,2,1,0,0,bird
|
||||
deer,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
dogfish,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,1,fish
|
||||
dolphin,0,0,0,1,0,1,1,1,1,1,0,1,0,1,0,1,mammal
|
||||
dove,0,1,1,0,1,0,0,0,1,1,0,0,2,1,1,0,bird
|
||||
duck,0,1,1,0,1,1,0,0,1,1,0,0,2,1,0,0,bird
|
||||
elephant,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
flamingo,0,1,1,0,1,0,0,0,1,1,0,0,2,1,0,1,bird
|
||||
flea,0,0,1,0,0,0,0,0,0,1,0,0,6,0,0,0,insect
|
||||
frog,0,0,1,0,0,1,1,1,1,1,0,0,4,0,0,0,amphibian
|
||||
frog,0,0,1,0,0,1,1,1,1,1,1,0,4,0,0,0,amphibian
|
||||
fruitbat,1,0,0,1,1,0,0,1,1,1,0,0,2,1,0,0,mammal
|
||||
giraffe,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
girl,1,0,0,1,0,0,1,1,1,1,0,0,2,0,1,1,mammal
|
||||
gnat,0,0,1,0,1,0,0,0,0,1,0,0,6,0,0,0,insect
|
||||
goat,1,0,0,1,0,0,0,1,1,1,0,0,4,1,1,1,mammal
|
||||
gorilla,1,0,0,1,0,0,0,1,1,1,0,0,2,0,0,1,mammal
|
||||
gull,0,1,1,0,1,1,1,0,1,1,0,0,2,1,0,0,bird
|
||||
haddock,0,0,1,0,0,1,0,1,1,0,0,1,0,1,0,0,fish
|
||||
hamster,1,0,0,1,0,0,0,1,1,1,0,0,4,1,1,0,mammal
|
||||
hare,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,0,mammal
|
||||
hawk,0,1,1,0,1,0,1,0,1,1,0,0,2,1,0,0,bird
|
||||
herring,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,0,fish
|
||||
honeybee,1,0,1,0,1,0,0,0,0,1,1,0,6,0,1,0,insect
|
||||
housefly,1,0,1,0,1,0,0,0,0,1,0,0,6,0,0,0,insect
|
||||
kiwi,0,1,1,0,0,0,1,0,1,1,0,0,2,1,0,0,bird
|
||||
ladybird,0,0,1,0,1,0,1,0,0,1,0,0,6,0,0,0,insect
|
||||
lark,0,1,1,0,1,0,0,0,1,1,0,0,2,1,0,0,bird
|
||||
leopard,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
lion,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
lobster,0,0,1,0,0,1,1,0,0,0,0,0,6,0,0,0,shellfish
|
||||
lynx,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
mink,1,0,0,1,0,1,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
mole,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,0,mammal
|
||||
mongoose,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
moth,1,0,1,0,1,0,0,0,0,1,0,0,6,0,0,0,insect
|
||||
newt,0,0,1,0,0,1,1,1,1,1,0,0,4,1,0,0,amphibian
|
||||
octopus,0,0,1,0,0,1,1,0,0,0,0,0,8,0,0,1,shellfish
|
||||
opossum,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,0,mammal
|
||||
oryx,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,1,mammal
|
||||
ostrich,0,1,1,0,0,0,0,0,1,1,0,0,2,1,0,1,bird
|
||||
parakeet,0,1,1,0,1,0,0,0,1,1,0,0,2,1,1,0,bird
|
||||
penguin,0,1,1,0,0,1,1,0,1,1,0,0,2,1,0,1,bird
|
||||
pheasant,0,1,1,0,1,0,0,0,1,1,0,0,2,1,0,0,bird
|
||||
pike,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,1,fish
|
||||
piranha,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,0,fish
|
||||
pitviper,0,0,1,0,0,0,1,1,1,1,1,0,0,1,0,0,reptile
|
||||
platypus,1,0,1,1,0,1,1,0,1,1,0,0,4,1,0,1,mammal
|
||||
polecat,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
pony,1,0,0,1,0,0,0,1,1,1,0,0,4,1,1,1,mammal
|
||||
porpoise,0,0,0,1,0,1,1,1,1,1,0,1,0,1,0,1,mammal
|
||||
puma,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
pussycat,1,0,0,1,0,0,1,1,1,1,0,0,4,1,1,1,mammal
|
||||
raccoon,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
reindeer,1,0,0,1,0,0,0,1,1,1,0,0,4,1,1,1,mammal
|
||||
rhea,0,1,1,0,0,0,1,0,1,1,0,0,2,1,0,1,bird
|
||||
scorpion,0,0,0,0,0,0,1,0,0,1,1,0,8,1,0,0,shellfish
|
||||
seahorse,0,0,1,0,0,1,0,1,1,0,0,1,0,1,0,0,fish
|
||||
seal,1,0,0,1,0,1,1,1,1,1,0,1,0,0,0,1,mammal
|
||||
sealion,1,0,0,1,0,1,1,1,1,1,0,1,2,1,0,1,mammal
|
||||
seasnake,0,0,0,0,0,1,1,1,1,0,1,0,0,1,0,0,reptile
|
||||
seawasp,0,0,1,0,0,1,1,0,0,0,1,0,0,0,0,0,shellfish
|
||||
skimmer,0,1,1,0,1,1,1,0,1,1,0,0,2,1,0,0,bird
|
||||
skua,0,1,1,0,1,1,1,0,1,1,0,0,2,1,0,0,bird
|
||||
slowworm,0,0,1,0,0,0,1,1,1,1,0,0,0,1,0,0,reptile
|
||||
slug,0,0,1,0,0,0,0,0,0,1,0,0,0,0,0,0,shellfish
|
||||
sole,0,0,1,0,0,1,0,1,1,0,0,1,0,1,0,0,fish
|
||||
sparrow,0,1,1,0,1,0,0,0,1,1,0,0,2,1,0,0,bird
|
||||
squirrel,1,0,0,1,0,0,0,1,1,1,0,0,2,1,0,0,mammal
|
||||
starfish,0,0,1,0,0,1,1,0,0,0,0,0,5,0,0,0,shellfish
|
||||
stingray,0,0,1,0,0,1,1,1,1,0,1,1,0,1,0,1,fish
|
||||
swan,0,1,1,0,1,1,0,0,1,1,0,0,2,1,0,1,bird
|
||||
termite,0,0,1,0,0,0,0,0,0,1,0,0,6,0,0,0,insect
|
||||
toad,0,0,1,0,0,1,0,1,1,1,0,0,4,0,0,0,amphibian
|
||||
tortoise,0,0,1,0,0,0,0,0,1,1,0,0,4,1,0,1,reptile
|
||||
tuatara,0,0,1,0,0,0,1,1,1,1,0,0,4,1,0,0,reptile
|
||||
tuna,0,0,1,0,0,1,1,1,1,0,0,1,0,1,0,1,fish
|
||||
vampire,1,0,0,1,1,0,0,1,1,1,0,0,2,1,0,0,mammal
|
||||
vole,1,0,0,1,0,0,0,1,1,1,0,0,4,1,0,0,mammal
|
||||
vulture,0,1,1,0,1,0,1,0,1,1,0,0,2,1,0,1,bird
|
||||
wallaby,1,0,0,1,0,0,0,1,1,1,0,0,2,1,0,1,mammal
|
||||
wasp,1,0,1,0,1,0,0,0,0,1,1,0,6,0,0,0,insect
|
||||
wolf,1,0,0,1,0,0,1,1,1,1,0,0,4,1,0,1,mammal
|
||||
worm,0,0,1,0,0,0,0,0,0,1,0,0,0,0,0,0,shellfish
|
||||
wren,0,1,1,0,1,0,0,0,1,1,0,0,2,1,0,0,bird
|
||||
|
@@ -0,0 +1,68 @@
|
||||
Source: http://www1.ics.uci.edu/pub/machine-learning-databases/zoo/',
|
||||
1. Title: Zoo database
|
||||
|
||||
2. Source Information
|
||||
-- Creator: Richard Forsyth
|
||||
-- Donor: Richard S. Forsyth
|
||||
8 Grosvenor Avenue
|
||||
Mapperley Park
|
||||
Nottingham NG3 5DX
|
||||
0602-621676
|
||||
-- Date: 5/15/1990
|
||||
|
||||
3. Past Usage:
|
||||
-- None known other than what is shown in Forsyth's PC/BEAGLE User's Guide.
|
||||
|
||||
4. Relevant Information:
|
||||
-- A simple database containing 17 Boolean-valued attributes. The "type"
|
||||
attribute appears to be the class attribute. Here is a breakdown of
|
||||
which animals are in which type: (I find it unusual that there are
|
||||
2 instances of "frog" and one of "girl"!)
|
||||
|
||||
Class# Set of animals:
|
||||
====== ===============================================================
|
||||
1 (41) aardvark, antelope, bear, boar, buffalo, calf,
|
||||
cavy, cheetah, deer, dolphin, elephant,
|
||||
fruitbat, giraffe, girl, goat, gorilla, hamster,
|
||||
hare, leopard, lion, lynx, mink, mole, mongoose,
|
||||
opossum, oryx, platypus, polecat, pony,
|
||||
porpoise, puma, pussycat, raccoon, reindeer,
|
||||
seal, sealion, squirrel, vampire, vole, wallaby,wolf
|
||||
2 (20) chicken, crow, dove, duck, flamingo, gull, hawk,
|
||||
kiwi, lark, ostrich, parakeet, penguin, pheasant,
|
||||
rhea, skimmer, skua, sparrow, swan, vulture, wren
|
||||
3 (5) pitviper, seasnake, slowworm, tortoise, tuatara
|
||||
4 (13) bass, carp, catfish, chub, dogfish, haddock,
|
||||
herring, pike, piranha, seahorse, sole, stingray, tuna
|
||||
5 (4) frog, frog, newt, toad
|
||||
6 (8) flea, gnat, honeybee, housefly, ladybird, moth, termite, wasp
|
||||
7 (10) clam, crab, crayfish, lobster, octopus,
|
||||
scorpion, seawasp, slug, starfish, worm
|
||||
|
||||
5. Number of Instances: 101
|
||||
|
||||
6. Number of Attributes: 18 (animal name, 15 Boolean attributes, 2 numerics)
|
||||
|
||||
7. Attribute Information: (name of attribute and type of value domain)
|
||||
1. animal name: Unique for each instance
|
||||
2. hair Boolean
|
||||
3. feathers Boolean
|
||||
4. eggs Boolean
|
||||
5. milk Boolean
|
||||
6. airborne Boolean
|
||||
7. aquatic Boolean
|
||||
8. predator Boolean
|
||||
9. toothed Boolean
|
||||
10. backbone Boolean
|
||||
11. breathes Boolean
|
||||
12. venomous Boolean
|
||||
13. fins Boolean
|
||||
14. legs Numeric (set of values: {0,2,4,5,6,8})
|
||||
15. tail Boolean
|
||||
16. domestic Boolean
|
||||
17. catsize Boolean
|
||||
18. type Numeric (integer values in range [1,7])
|
||||
|
||||
8. Missing Attribute Values: None
|
||||
|
||||
9. Class Distribution: Given above
|
||||
Whitespace-only changes.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,183 @@
|
||||
import sys
|
||||
sys.path.insert(1, '../')
|
||||
from utils.utils import *
|
||||
from collections import defaultdict
|
||||
|
||||
|
||||
### the ABC of data set
|
||||
class DataSet:
|
||||
"""
|
||||
A data set for a machine learning problem. It has the following fields:
|
||||
|
||||
d.examples A list of examples. Each one is a list of attribute values.
|
||||
d.attrs A list of integers to index into an example, so example[attr]
|
||||
gives a value. Normally the same as range(len(d.examples[0])).
|
||||
d.attr_names Optional list of mnemonic names for corresponding attrs.
|
||||
d.target The attribute that a learning algorithm will try to predict.
|
||||
By default the final attribute.
|
||||
d.inputs The list of attrs without the target.
|
||||
d.values A list of lists: each sublist is the set of possible
|
||||
values for the corresponding attribute. If initially None,
|
||||
it is computed from the known examples by self.set_problem.
|
||||
If not None, an erroneous value raises ValueError.
|
||||
d.distance A function from a pair of examples to a non-negative number.
|
||||
Should be symmetric, etc. Defaults to mean_boolean_error
|
||||
since that can handle any field types.
|
||||
d.name Name of the data set (for output display only).
|
||||
d.source URL or other source where the data came from.
|
||||
d.exclude A list of attribute indexes to exclude from d.inputs. Elements
|
||||
of this list can either be integers (attrs) or attr_names.
|
||||
|
||||
Normally, you call the constructor and you're done; then you just
|
||||
access fields like d.examples and d.target and d.inputs.
|
||||
"""
|
||||
|
||||
def __init__(self, examples=None, attrs=None, attr_names=None, target=-1, inputs=None,
|
||||
values=None, distance=mean_boolean_error, name='', source='', exclude=()):
|
||||
"""
|
||||
Accepts any of DataSet's fields. Examples can also be a
|
||||
string or file from which to parse examples using parse_csv.
|
||||
Optional parameter: exclude, as documented in .set_problem().
|
||||
>>> DataSet(examples='1, 2, 3')
|
||||
<DataSet(): 1 examples, 3 attributes>
|
||||
"""
|
||||
self.name = name
|
||||
self.source = source
|
||||
self.values = values
|
||||
self.distance = distance
|
||||
self.got_values_flag = bool(values)
|
||||
|
||||
# initialize .examples from string or list or data directory
|
||||
if isinstance(examples, str):
|
||||
self.examples = parse_csv(examples)
|
||||
elif examples is None:
|
||||
self.examples = parse_csv(open_data(name + '.csv').read())
|
||||
else:
|
||||
self.examples = examples
|
||||
|
||||
# attrs are the indices of examples, unless otherwise stated.
|
||||
if self.examples is not None and attrs is None:
|
||||
attrs = list(range(len(self.examples[0])))
|
||||
|
||||
self.attrs = attrs
|
||||
|
||||
# initialize .attr_names from string, list, or by default
|
||||
if isinstance(attr_names, str):
|
||||
self.attr_names = attr_names.split()
|
||||
else:
|
||||
self.attr_names = attr_names or attrs
|
||||
self.set_problem(target, inputs=inputs, exclude=exclude)
|
||||
|
||||
def set_problem(self, target, inputs=None, exclude=()):
|
||||
"""
|
||||
Set (or change) the target and/or inputs.
|
||||
This way, one DataSet can be used multiple ways. inputs, if specified,
|
||||
is a list of attributes, or specify exclude as a list of attributes
|
||||
to not use in inputs. Attributes can be -n .. n, or an attr_name.
|
||||
Also computes the list of possible values, if that wasn't done yet.
|
||||
"""
|
||||
self.target = self.attr_num(target)
|
||||
exclude = list(map(self.attr_num, exclude))
|
||||
if inputs:
|
||||
self.inputs = remove_all(self.target, inputs)
|
||||
else:
|
||||
self.inputs = [a for a in self.attrs if a != self.target and a not in exclude]
|
||||
if not self.values:
|
||||
self.update_values()
|
||||
self.check_me()
|
||||
|
||||
def check_me(self):
|
||||
"""Check that my fields make sense."""
|
||||
assert len(self.attr_names) == len(self.attrs)
|
||||
assert self.target in self.attrs
|
||||
assert self.target not in self.inputs
|
||||
assert set(self.inputs).issubset(set(self.attrs))
|
||||
if self.got_values_flag:
|
||||
# only check if values are provided while initializing DataSet
|
||||
list(map(self.check_example, self.examples))
|
||||
|
||||
def add_example(self, example):
|
||||
"""Add an example to the list of examples, checking it first."""
|
||||
self.check_example(example)
|
||||
self.examples.append(example)
|
||||
|
||||
def check_example(self, example):
|
||||
"""Raise ValueError if example has any invalid values."""
|
||||
if self.values:
|
||||
for a in self.attrs:
|
||||
if example[a] not in self.values[a]:
|
||||
raise ValueError('Bad value {} for attribute {} in {}'
|
||||
.format(example[a], self.attr_names[a], example))
|
||||
|
||||
def attr_num(self, attr):
|
||||
"""Returns the number used for attr, which can be a name, or -n .. n-1."""
|
||||
if isinstance(attr, str):
|
||||
return self.attr_names.index(attr)
|
||||
elif attr < 0:
|
||||
return len(self.attrs) + attr
|
||||
else:
|
||||
return attr
|
||||
|
||||
def update_values(self):
|
||||
self.values = list(map(unique, zip(*self.examples)))
|
||||
|
||||
def sanitize(self, example):
|
||||
"""Return a copy of example, with non-input attributes replaced by None."""
|
||||
return [attr_i if i in self.inputs else None for i, attr_i in enumerate(example)][:-1]
|
||||
|
||||
def classes_to_numbers(self, classes=None):
|
||||
"""Converts class names to numbers."""
|
||||
if not classes:
|
||||
# if classes were not given, extract them from values
|
||||
classes = sorted(self.values[self.target])
|
||||
for item in self.examples:
|
||||
item[self.target] = classes.index(item[self.target])
|
||||
|
||||
def remove_examples(self, value=''):
|
||||
"""Remove examples that contain given value."""
|
||||
self.examples = [x for x in self.examples if value not in x]
|
||||
self.update_values()
|
||||
|
||||
def split_values_by_classes(self):
|
||||
"""Split values into buckets according to their class."""
|
||||
buckets = defaultdict(lambda: [])
|
||||
target_names = self.values[self.target]
|
||||
|
||||
for v in self.examples:
|
||||
item = [a for a in v if a not in target_names] # remove target from item
|
||||
buckets[v[self.target]].append(item) # add item to bucket of its class
|
||||
|
||||
return buckets
|
||||
|
||||
def find_means_and_deviations(self):
|
||||
"""
|
||||
Finds the means and standard deviations of self.dataset.
|
||||
means : a dictionary for each class/target. Holds a list of the means
|
||||
of the features for the class.
|
||||
deviations: a dictionary for each class/target. Holds a list of the sample
|
||||
standard deviations of the features for the class.
|
||||
"""
|
||||
target_names = self.values[self.target]
|
||||
feature_numbers = len(self.inputs)
|
||||
|
||||
item_buckets = self.split_values_by_classes()
|
||||
|
||||
means = defaultdict(lambda: [0] * feature_numbers)
|
||||
deviations = defaultdict(lambda: [0] * feature_numbers)
|
||||
|
||||
for t in target_names:
|
||||
# find all the item feature values for item in class t
|
||||
features = [[] for _ in range(feature_numbers)]
|
||||
for item in item_buckets[t]:
|
||||
for i in range(feature_numbers):
|
||||
features[i].append(item[i])
|
||||
|
||||
# calculate means and deviations fo the class
|
||||
for i in range(feature_numbers):
|
||||
means[t][i] = mean(features[i])
|
||||
deviations[t][i] = stdev(features[i])
|
||||
|
||||
return means, deviations
|
||||
|
||||
def __repr__(self):
|
||||
return '<DataSet({}): {:d} examples, {:d} attributes>'.format(self.name, len(self.examples), len(self.attrs))
|
||||
+198
@@ -0,0 +1,198 @@
|
||||
### system related operations
|
||||
import os
|
||||
from inspect import getsource
|
||||
from IPython.display import HTML
|
||||
from IPython.display import display
|
||||
|
||||
|
||||
def psource(*functions):
|
||||
"""Print the source code for the given function(s)."""
|
||||
source_code = '\n\n'.join(getsource(fn) for fn in functions)
|
||||
try:
|
||||
from pygments.formatters import HtmlFormatter
|
||||
from pygments.lexers import PythonLexer
|
||||
from pygments import highlight
|
||||
|
||||
display(HTML(highlight(source_code, PythonLexer(), HtmlFormatter(full=True))))
|
||||
|
||||
except ImportError:
|
||||
print(source_code)
|
||||
|
||||
|
||||
def parse_csv(input, delim=','):
|
||||
r"""
|
||||
Input is a string consisting of lines, each line has comma-delimited
|
||||
fields. Convert this into a list of lists. Blank lines are skipped.
|
||||
Fields that look like numbers are converted to numbers.
|
||||
The delim defaults to ',' but '\t' and None are also reasonable values.
|
||||
>>> parse_csv('1, 2, 3 \n 0, 2, na')
|
||||
[[1, 2, 3], [0, 2, 'na']]
|
||||
"""
|
||||
lines = [line for line in input.splitlines() if line.strip()]
|
||||
return [list(map(num_or_str, line.split(delim))) for line in lines]
|
||||
|
||||
|
||||
def open_data(name, mode='r'):
|
||||
data_root = os.path.dirname(__file__)
|
||||
data_file = os.path.join(data_root, *[os.pardir, 'data', name])
|
||||
|
||||
return open(data_file, mode=mode)
|
||||
|
||||
|
||||
### math and data structure
|
||||
import numpy as np
|
||||
import math
|
||||
from statistics import mean, stdev
|
||||
import random
|
||||
|
||||
|
||||
def normalize(dist):
|
||||
"""Multiply each number by a constant such that the sum is 1.0"""
|
||||
if isinstance(dist, dict):
|
||||
total = sum(dist.values())
|
||||
for key in dist:
|
||||
dist[key] = dist[key] / total
|
||||
assert 0 <= dist[key] <= 1 # probabilities must be between 0 and 1
|
||||
return dist
|
||||
total = sum(dist)
|
||||
return [(n / total) for n in dist]
|
||||
|
||||
|
||||
def random_weights(min_value, max_value, num_weights):
|
||||
return [random.uniform(min_value, max_value) for _ in range(num_weights)]
|
||||
|
||||
|
||||
def sigmoid(x):
|
||||
"""Return activation value of x with sigmoid function."""
|
||||
return 1 / (1 + np.exp(-x))
|
||||
|
||||
|
||||
def sigmoid_derivative(value):
|
||||
return value * (1 - value)
|
||||
|
||||
|
||||
def remove_all(item, seq):
|
||||
"""Return a copy of seq (or string) with all occurrences of item removed."""
|
||||
if isinstance(seq, str):
|
||||
return seq.replace(item, '')
|
||||
elif isinstance(seq, set):
|
||||
rest = seq.copy()
|
||||
rest.remove(item)
|
||||
return rest
|
||||
else:
|
||||
return [x for x in seq if x != item]
|
||||
|
||||
|
||||
def unique(seq):
|
||||
"""Remove duplicate elements from seq. Assumes hashable elements."""
|
||||
return list(set(seq))
|
||||
|
||||
|
||||
def num_or_str(x): # TODO: rename as `atom`
|
||||
"""The argument is a string; convert to a number if
|
||||
possible, or strip it."""
|
||||
try:
|
||||
return int(x)
|
||||
except ValueError:
|
||||
try:
|
||||
return float(x)
|
||||
except ValueError:
|
||||
return str(x).strip()
|
||||
|
||||
|
||||
def euclidean_distance(x, y):
|
||||
return np.sqrt(sum((_x - _y) ** 2 for _x, _y in zip(x, y)))
|
||||
|
||||
|
||||
def manhattan_distance(x, y):
|
||||
return sum(abs(_x - _y) for _x, _y in zip(x, y))
|
||||
|
||||
|
||||
def hamming_distance(x, y):
|
||||
return sum(_x != _y for _x, _y in zip(x, y))
|
||||
|
||||
|
||||
def rms_error(x, y):
|
||||
return np.sqrt(ms_error(x, y))
|
||||
|
||||
|
||||
def ms_error(x, y):
|
||||
return mean((x - y) ** 2 for x, y in zip(x, y))
|
||||
|
||||
|
||||
def mean_error(x, y):
|
||||
return mean(abs(x - y) for x, y in zip(x, y))
|
||||
|
||||
|
||||
def mean_boolean_error(x, y):
|
||||
return mean(_x != _y for _x, _y in zip(x, y))
|
||||
|
||||
|
||||
identity = lambda x: x
|
||||
|
||||
|
||||
def argmin_random_tie(seq, key=identity):
|
||||
"""Return a minimum element of seq; break ties at random."""
|
||||
return min(shuffled(seq), key=key)
|
||||
|
||||
|
||||
def argmax_random_tie(seq, key=identity):
|
||||
"""Return an element with highest fn(seq[i]) score; break ties at random."""
|
||||
return max(shuffled(seq), key=key)
|
||||
|
||||
|
||||
def shuffled(iterable):
|
||||
"""Randomly shuffle a copy of iterable."""
|
||||
items = list(iterable)
|
||||
random.shuffle(items)
|
||||
return items
|
||||
|
||||
|
||||
### learner related operations
|
||||
class Activation:
|
||||
|
||||
def function(self, x):
|
||||
return NotImplementedError
|
||||
|
||||
def derivative(self, x):
|
||||
return NotImplementedError
|
||||
|
||||
def __call__(self, x):
|
||||
return self.function(x)
|
||||
|
||||
|
||||
class Sigmoid(Activation):
|
||||
|
||||
def function(self, x):
|
||||
return 1 / (1 + np.exp(-x))
|
||||
|
||||
def derivative(self, value):
|
||||
return value * (1 - value)
|
||||
|
||||
|
||||
def err_ratio(learner, dataset, examples=None):
|
||||
"""
|
||||
Return the proportion of the examples that are NOT correctly predicted.
|
||||
verbose - 0: No output; 1: Output wrong; 2 (or greater): Output correct
|
||||
"""
|
||||
examples = examples or dataset.examples
|
||||
if len(examples) == 0:
|
||||
return 0.0
|
||||
right = 0
|
||||
for example in examples:
|
||||
desired = example[dataset.target]
|
||||
output = learner.predict(dataset.sanitize(example))
|
||||
if output == desired:
|
||||
right += 1
|
||||
return 1 - (right / len(examples))
|
||||
|
||||
|
||||
def grade_learner(learner, tests):
|
||||
"""
|
||||
Grades the given learner based on how many tests it passes.
|
||||
tests is a list with each element in the form: (values, output).
|
||||
"""
|
||||
# for X, y in tests:
|
||||
# print(learner.predict(X), y)
|
||||
return mean([int(learner.predict(X) == y) for X, y in tests])
|
||||
|
||||
Reference in new issue
Block a user