<?xml version="1.0" encoding="utf-8"?>
<feed xmlns="http://www.w3.org/2005/Atom"><title>Data and AI Engineering</title>
<generator>Emacs webfeeder.el</generator>
<link href="https://dataai.ng"/>
<link href="https://dataai.ng/atom.xml" rel="self"/>
<id>https://dataai.ng/atom.xml</id>
<updated>2026-07-31T15:09:54+03:00</updated>
<entry>
  <title>Implementing Search for Swahili Text</title>
  <author><name>KE Programmer</name></author>
  <summary>Published on Aug 20, 2024 14:57 by KE Programmer</summary>
  <content type="html"><![CDATA[<div id="content" class="content">
 <h1 class="title">Implementing Search for Swahili Text
 <br></br> <span class="subtitle">Published on Aug 20, 2024 14:57 by KE Programmer</span>
</h1>
 <div id="table-of-contents" role="doc-toc">
 <h2>Table of Contents</h2>
 <div id="text-table-of-contents" role="doc-toc">
 <ul> <li> <a href="#org6b432df">1. Introduction</a></li>
 <li> <a href="#orgc754d83">2. Setup</a></li>
 <li> <a href="#org96dc887">3. Data Processing</a>
 <ul> <li> <a href="#org189cc14">3.1. Get the Data</a></li>
 <li> <a href="#org2255be3">3.2. Load the Data</a></li>
 <li> <a href="#org4e12b40">3.3. Explore the Data</a></li>
</ul></li>
 <li> <a href="#orgdcc5542">4. Implementing Text Search</a>
 <ul> <li> <a href="#org28bcbe0">4.1. Keyword Filtering</a></li>
 <li> <a href="#orge72d7eb">4.2. Vectorization</a>
 <ul> <li> <a href="#org9ab2e93">4.2.1. Count Vectorization</a></li>
 <li> <a href="#org6255a02">4.2.2. TF-IDF Vectorization</a></li>
</ul></li>
 <li> <a href="#orgb3cd5e4">4.3. Embeddings</a>
 <ul> <li> <a href="#org914ce40">4.3.1. Singular Value Decomposition (SVD)</a></li>
 <li> <a href="#org2166ff3">4.3.2. BERT</a></li>
</ul></li>
</ul></li>
 <li> <a href="#org00b0096">5. Conclusion</a></li>
</ul></div>
</div>
 <div id="outline-container-org6b432df" class="outline-2">
 <h2 id="org6b432df"> <span class="section-number-2">1.</span> Introduction</h2>
 <div class="outline-text-2" id="text-1">
 <p>
A treatment of various search techniques for Swahili text, from simple
keyword filtering to more sophisticated semantic search. We'll use
data from the  <a href="https://github.com/masakhane-io/masakhane-ner/tree/main">MasakhaNER</a> project for demonstration.
</p>
</div>
</div>
 <div id="outline-container-orgc754d83" class="outline-2">
 <h2 id="orgc754d83"> <span class="section-number-2">2.</span> Setup</h2>
 <div class="outline-text-2" id="text-2">
 <p>
First step is to setup the dependencies and import commonly used modules:
</p>

 <div class="org-src-container">
 <pre class="src src-sh">pip install -q requests pandas scikit-learn jupyter transformers tqdm
pip install -q torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">import</span> pandas  <span style="font-weight: bold;">as</span> pd
 <span style="font-weight: bold;">import</span> requests
 <span style="font-weight: bold;">import</span> io
 <span style="font-weight: bold;">import</span> numpy  <span style="font-weight: bold;">as</span> np
</pre>
</div>
</div>
</div>
 <div id="outline-container-org96dc887" class="outline-2">
 <h2 id="org96dc887"> <span class="section-number-2">3.</span> Data Processing</h2>
 <div class="outline-text-2" id="text-3">
</div>
 <div id="outline-container-org189cc14" class="outline-3">
 <h3 id="org189cc14"> <span class="section-number-3">3.1.</span> Get the Data</h3>
 <div class="outline-text-3" id="text-3-1">
 <p>
The data we will be using is Swahili News text data from the
 <a href="https://github.com/masakhane-io/masakhane-ner/tree/main/text_by_language/swahili">masakhane-ner</a> repository, whose origin is the Swahili version of
 <a href="https://www.voaswahili.com/z/2772">Voice of America</a>.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">docs_url</span> =  <span style="font-style: italic;">'https://raw.githubusercontent.com/masakhane-io/masakhane-ner/main/text_by_language/swahili/voa_clean_final.txt'</span>
 <span style="font-weight: bold; font-style: italic;">docs_response</span> = requests.get(docs_url)
 <span style="font-weight: bold; font-style: italic;">documents_raw</span> = docs_response.text
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">documents_raw[:999]
</pre>
</div>

 <pre class="example">
Wizara ya afya ya Tanzania imeripoti Jumatatu kuwa, watu takriban 14 zaidi wamepata maambukizi ya Covid-19.
Walioambukizwa wote ni raia wa Tanzania, 13 wakiwa Dar-es-salaam na mmoja mjini Arusha.
Wizara ya afya imeripoti kwamba juhudi za kufuatilia watu waliokuwa karibu na wagonjwa zinaendelea.
Wakati wa maadimisho ya pasaka, wakristo walikusanyika kanisani kwa maombi bila kuzingatia ushauri wa wataalam wa afya.
Kuna mijadala kwenye mitandao ya kijamii Tanzania, kuhusu hatua zinazochukuliwa kudhibithi maambukizi nchini humo.
Nchini Afrika kusini watu 145 zaidi, wameambukizwa virusi vya Corona, na kujumulisha idadi ya watu 2,173 ambao wameambukizwa virusi vya Corona nchini humo.
Taarifa ya wizara ya afya hata hivyo haijasema idadi ya watu ambao wamekufa wala kupona kutokana na virusi vya Corona nchini humo.
Nchini Sudan, maafisa wameongeza mikakati zaidi ya kuzuia virusi vya Corona kusambaa.
Wamepiga marufuku usafiri wa magari kati ya miji na kutekeleza sheria za hali ya dharura ili ku
</pre>
</div>
</div>
 <div id="outline-container-org2255be3" class="outline-3">
 <h3 id="org2255be3"> <span class="section-number-3">3.2.</span> Load the Data</h3>
 <div class="outline-text-3" id="text-3-2">
 <p>
The data is composed of randomly ordered sentences from the News
sources.  We'll treat each sentence as a single document in our
corpus. We extract the sentences line by line into a pandas dataframe.
</p>

 <div class="org-src-container">
 <pre class="src src-python">pd.set_option( <span style="font-style: italic;">'display.max_colwidth'</span>, 999)  <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">avoid truncation of the column
</span>
 <span style="font-weight: bold; font-style: italic;">in_memory_file</span> = io.StringIO(documents_raw)
 <span style="font-weight: bold; font-style: italic;">df</span> = pd.DataFrame([l.strip()  <span style="font-weight: bold;">for</span> l  <span style="font-weight: bold;">in</span> in_memory_file], columns=[ <span style="font-style: italic;">'documents'</span>])
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">df.tail()
</pre>
</div>

 <pre class="example">
                                                                                                                                                                                                                                                                                              documents
7667                                                                   Wakati huohuo upande wa utetezi uliomwakilisha Jaji umesema hauna tena sababu ya kuwaita mashahidi wake, lakini itaomba mahakama hiyo kutupia mbali kesi hiyo kwa sababu serikali ya Nigeria imeshindwa kuthibitisha madai yake.
7668                                                                                                            Baadhi ya wananchi wa Nigeria wanadai kuwa hatua iliyochukuliwa na Rais Buhari, ni njama ya kumweka mwengine kutoka upande wa Kaskazini mwa Nigeria kuchukuwa cheo hicho cha Jaji Mkuu.
7669  Madai yao pia yanalenga suala la njama ya kurudi tena madarakani wakidai kuwa Buhari alikuwa akijitayarisha kushindania urais awamu ya pili na iwapo angeshindwa kesi ikifikishwa Mahakama Kuu kabisa, atakuwa na mtu wake wa karibu atakaye mwonea huruma na kuhakikisha kwamba anapata ushindi.
7670                                                                                                                                                                                                                                        Imetayarishwa na Mwandishi wetu, Collins Atohengbe, Nigeria
7671                                   Walinzi wa pwani ya Libya wamekamata wahamiaji 400 waliokuwa wakonjiani katika pwani ya Mediterranean ya nchi hiyo wakielekea Ulaya na kuwarejesha katika mji mkuu wa Tripoli masaa 24 yaliyopita, Shirika la uhamiaji la Umoja wa Mataifa UN limesema Jumapili.
</pre>


 <div class="org-src-container">
 <pre class="src src-python">df.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">7672</td>
 <td class="org-right">1</td>
</tr></tbody></table> <p>
We have 7672 documents in total.
</p>
</div>
</div>
 <div id="outline-container-org4e12b40" class="outline-3">
 <h3 id="org4e12b40"> <span class="section-number-3">3.3.</span> Explore the Data</h3>
 <div class="outline-text-3" id="text-3-3">
 <p>
The data includes text from 2020 when the COVID-19 pandemic was a
major news item. Let's say that from our dataset we want to find
documents related to Africa's response to the pandemic. We'll
use the Swahili query below:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">query</span> =  <span style="font-style: italic;">"nchi za afrika zinajadiliana na shirika la afya kutafuta njia za kukabiliana na janga la corona"</span>
</pre>
</div>
</div>
</div>
</div>
 <div id="outline-container-orgdcc5542" class="outline-2">
 <h2 id="orgdcc5542"> <span class="section-number-2">4.</span> Implementing Text Search</h2>
 <div class="outline-text-2" id="text-4">
</div>
 <div id="outline-container-org28bcbe0" class="outline-3">
 <h3 id="org28bcbe0"> <span class="section-number-3">4.1.</span> Keyword Filtering</h3>
 <div class="outline-text-3" id="text-4-1">
 <p>
A simple technique is to use keywords from our query to filter documents
that may have the information we require. We only match documents that
contain only the keywords we've selected.
</p>

 <p>
In the example below, we create filters for rows containing each word
individually then combine the filters to filter in only rows with all
the words.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">keywords</span> = [ <span style="font-style: italic;">'afrika'</span>,  <span style="font-style: italic;">'afya'</span>,  <span style="font-style: italic;">'corona'</span>]
 <span style="font-weight: bold; font-style: italic;">filters</span> = [df.documents. <span style="font-weight: bold;">str</span>.contains(word,  <span style="font-weight: bold;">case</span>= <span style="font-weight: bold; text-decoration: underline;">False</span>)  <span style="font-weight: bold;">for</span> word  <span style="font-weight: bold;">in</span> keywords]
 <span style="font-weight: bold; font-style: italic;">combined_filter</span> = np.vstack(filters). <span style="font-weight: bold;">all</span>(axis=0)
df[combined_filter]
</pre>
</div>

 <pre class="example">
                                                                                                                                                                                                                                                                                                                        documents
1706  Waziri wa Afya Zweli Mkhize amethibitisha vifo hivyo, akisema marehemu wote hao walifia Magharibi mwa Cape province  Afrika Kusini ina zaidi ya watu 1,000 walioambukizwa virusi hivyo, ikiwa idadi kubwa zaidi katika Afrika, ambapo zaidi ya watu 3,000 barani humo wamethibitishwa kuwa na ugonjwa wa virusi vya corona.
</pre>


 <p>
Only one document matches all the selected keywords.
</p>

 <p>
While this method is straightforward, it is sensitive to the
combination of keywords selected and possible misspellings making it
useful for mostly basic searches. When there's more than one match,
the results have to be examined further to determine the most relevant
ones.
</p>
</div>
</div>
 <div id="outline-container-orge72d7eb" class="outline-3">
 <h3 id="orge72d7eb"> <span class="section-number-3">4.2.</span> Vectorization</h3>
 <div class="outline-text-3" id="text-4-2">
 <p>
In vectorization, we move from the textual representation to a
numerical one, which can help with ranking the results.
</p>

 <p>
All the unique words in the document corpus are identified, ordered
and assigned a unique index as their identity within the
vocabulary. The collection of documents will be represented as a table
where each row represents a single document and each column represents
a word in the derived vocabulary. This is known as a  <b>document-term
matrix</b>.
</p>

 <p>
There are two types of vectorization we'll look at, Count
Vectorization and TF-IDF.
</p>
</div>
 <div id="outline-container-org9ab2e93" class="outline-4">
 <h4 id="org9ab2e93"> <span class="section-number-4">4.2.1.</span> Count Vectorization</h4>
 <div class="outline-text-4" id="text-4-2-1">
 <p>
In count vectorization, the value of each cell in the document-term
matrix represents the number of times the particular word appears in
the document.  Calling  <code>fit_transform</code> on the count vectorizer below
first creates the vocabulary by standardizing and ordering all the
available terms, assigning a unique index to each (fit), then converts
each document into a row of word counts (transform).
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.feature_extraction.text  <span style="font-weight: bold;">import</span> CountVectorizer

 <span style="font-weight: bold; font-style: italic;">cv</span> = CountVectorizer()
 <span style="font-weight: bold; font-style: italic;">X</span> = cv.fit_transform(df.documents)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">X.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">7672</td>
 <td class="org-right">19056</td>
</tr></tbody></table> <p>
The resulting matrix has a number of rows equivalent to the number of
documents, and the unique terms extracted are 19056.
</p>

 <p>
We can inspect the last few words of the vocabulary:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">names</span> = cv.get_feature_names_out()
 <span style="font-weight: bold;">len</span>(names), names[-7:]
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-right">19056</td>
 <td class="org-left">array</td>
 <td class="org-left">((zulu zuma zungumza zuri zusha zweli évariste) dtype=object)</td>
</tr></tbody></table> <p>
Create a new dataframe for the document-term matrix:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">df_docs</span> = pd.DataFrame(X.toarray(), columns=names)
df_docs.head()
</pre>
</div>

 <pre class="example">
   00  000  002  01  02  023  041  043  ...  zuio  zulu  zuma  zungumza  zuri  zusha  zweli  évariste
0   0    0    0   0   0    0    0    0  ...     0     0     0         0     0      0      0         0
1   0    0    0   0   0    0    0    0  ...     0     0     0         0     0      0      0         0
2   0    0    0   0   0    0    0    0  ...     0     0     0         0     0      0      0         0
3   0    0    0   0   0    0    0    0  ...     0     0     0         0     0      0      0         0
4   0    0    0   0   0    0    0    0  ...     0     0     0         0     0      0      0         0

[5 rows x 19056 columns]
</pre>


 <p>
The document-term matrix is a sparse matrix because most words do not
occur in most documents, therefore most of the counts will be zero.
</p>

 <p>
Looking into the first document, we can filter out the zero values to
see the count of words that exist in it.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">first_doc</span> = df_docs.loc[0]
first_doc[first_doc > 0]
</pre>
</div>

 <pre class="example" id="org0af5c2b">
14            1
19            1
afya          1
covid         1
imeripoti     1
jumatatu      1
kuwa          1
maambukizi    1
takriban      1
tanzania      1
wamepata      1
watu          1
wizara        1
ya            3
zaidi         1
Name: 0, dtype: int64
</pre>

 <p>
The full document:
</p>

 <div class="org-src-container">
 <pre class="src src-python">df.loc[0]
</pre>
</div>

 <pre class="example">
documents    Wizara ya afya ya Tanzania imeripoti Jumatatu kuwa, watu takriban 14 zaidi wamepata maambukizi ya Covid-19.
Name: 0, dtype: str
</pre>
</div>
 <ol class="org-ol"> <li> <a id="org3358994"></a>Query-Document similarity <br></br> <div class="outline-text-5" id="text-4-2-1-1">
 <p>
To search through the documents using the query, we first need to
transform the query using the same vectorizer as the documents. This
maps the query to the same vector space; the length of the resulting
vector matches the vocabulary size and each index in it contains the
count of the specific term in the query.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">q</span> = cv.transform([query])
 <span style="font-weight: bold; font-style: italic;">df_query</span> = pd.DataFrame(q.toarray(), columns=names)
df_query.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">1</td>
 <td class="org-right">19056</td>
</tr></tbody></table> <p>
Most values are expected to be zeros as well.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">encoded_query</span> = df_query.loc[0]
encoded_query[encoded_query > 0]
</pre>
</div>

 <pre class="example" id="orgba4dca9">
afrika         1
afya           1
corona         1
janga          1
kukabiliana    1
kutafuta       1
la             2
na             2
nchi           1
njia           1
shirika        1
za             2
Name: 0, dtype: int64
</pre>

 <p>
The more the common words and counts between query and document, the more
similar they are. We use  <a href="https://www.mathsisfun.com/algebra/vectors-dot-product.html">dot product</a> to calculate the score between each
document and query, then rank by score. The closer the vectors of the
query and a particular document in the vector space are, the higher the
dot-product.
</p>

 <p>
For example, dot product between query and first document:
</p>

 <div class="org-src-container">
 <pre class="src src-python">(first_doc * encoded_query). <span style="font-weight: bold;">sum</span>()
</pre>
</div>

 <pre class="example">
1
</pre>


 <p>
Across all the docs:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">query_vector</span> = q.toarray().flatten()
 <span style="font-weight: bold; font-style: italic;">score</span> = X.dot(query_vector)
score.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">7672</td>
</tr></tbody></table> <p>
The score vector has the resulting dot products for each document. We
identify the highest one:
</p>

 <div class="org-src-container">
 <pre class="src src-python">score.argmax(), score[score.argmax()]
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">np.int64</td>
 <td class="org-left">(2433)</td>
 <td class="org-left">np.int64</td>
 <td class="org-left">(30)</td>
</tr></tbody></table> <p>
The highest match is document at index 2433, with a score of 30.
</p>

 <p>
The particular document from the dataframe.
</p>

 <div class="org-src-container">
 <pre class="src src-python">df.documents[2433]
</pre>
</div>

 <pre class="example">
”  Swali la msingi la Mueller  Swali la msingi la Mueller, mkurugenzi wa zamani wa FBI, ambalo anatafuta majibu : Ni iwapo Trump na wasaidizi wake walishirikiana na Warusi kuchafua kampeni ya mgombea wa chama cha Demokrat Hillary Clinton, mwaka 2016, kwa kutuma barua pepe zenye kudhalilisha zilizoibiwa kutoka Kamati ya Taifa ya chama cha Demokrat na mwenyekiti wa kampeni ya Clinton? Au iwapo Trump alikuwa amenufaika bila ya kukusudia na mbinu chafu za Russia? Na iwapo rais alijaribu kuharibu uchunguzi uliofuatia ili kujilinda yeye mwenyewe na washauri wa kisiasa na wasaidizi wake?  Huu ndio ujumbe wa Idara ya Sheria kwa bunge la Congress juu ya hitimisho la uchunguzi uliofanywa na Mueller.
</pre>


 <p>
We can show the top results ranked from highest to lowest score:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">top_idx</span> = np.argsort(-score)[:10]
df.iloc[top_idx]
</pre>
</div>

 <pre class="example" id="org8998ff1">
                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                       documents
2433  ”  Swali la msingi la Mueller  Swali la msingi la Mueller, mkurugenzi wa zamani wa FBI, ambalo anatafuta majibu : Ni iwapo Trump na wasaidizi wake walishirikiana na Warusi kuchafua kampeni ya mgombea wa chama cha Demokrat Hillary Clinton, mwaka 2016, kwa kutuma barua pepe zenye kudhalilisha zilizoibiwa kutoka Kamati ya Taifa ya chama cha Demokrat na mwenyekiti wa kampeni ya Clinton? Au iwapo Trump alikuwa amenufaika bila ya kukusudia na mbinu chafu za Russia? Na iwapo rais alijaribu kuharibu uchunguzi uliofuatia ili kujilinda yeye mwenyewe na washauri wa kisiasa na wasaidizi wake?  Huu ndio ujumbe wa Idara ya Sheria kwa bunge la Congress juu ya hitimisho la uchunguzi uliofanywa na Mueller.
6107                                                                                                                                                                                                                                                                                                                                      Lissu ameongeza kwamba “Kwa kutumia njia hizo za amri za kisiasa, vyombo vya ulinzi na usalama vikiwemo jeshi la wananchi la Tanzania, jeshi la polisi, taasisi ya kuzuia na kupambana na rushwa, na idara ya usalama wa taifa pamoja na mamlaka ya kodi TRA, vinatumika kukamata kwa nguvu na kutaifisha mali za wafanyabiashara wetu wa ndani na makampuni ya wawekezaji kutoka nje.
968                                                                                                                                                                                                                                Akizungumza na shirika la habari la AFP Akram Taher mumini moja amesema “Sherehe za Eid hazifani manmo wakati huu wa hali ya janga la corona – watu wanahisia ya kua na hofu”  Bara la Asia  Waislam katika bara la Asia – kutoka Indonesia hadi Pakistan, Malaysia na Afghanistan – wamekusanyika katika masoko katika shamrashamra za manunuzi ya sikukuu, wakikiuka muongozo wa kudhibiti virusi vya corona na wakati mwengine polisi wakijaribu kutawanya mikusanyiko ya makundi makubwa.
530                                                                                                                                                                                                                                                                                                                                        Changamoto za uokozi  Msemaji wa shirika la kimataifa la msalaba mwekundu Caroline Haga ameliambia shirika la habari la AFP kwamba wafanyakazi wa uokozi wanakabiliwa na changamoto kubwa kuwafikia wanaohitaji msaada na wana wasiwasi na huzuni, kwani siku ya Jumanne walifanikiwa kuwaokoa watu 167 na kwamba muda unakimbia haraka sana na watu bado wamekwama na wako hatarini.
842                                                                                                                                                                                      “Wakati hatuamini kuwa katika hatua hii, hali hiyo inahitaji kupitishwa azimio, kuna dalili zote za kututahadharisha kuwa mgogoro wa kuminywa kwa haki za binadamu unafukuta,” imesema barua hiyo, ambayo imesainiwa na Mtandao wa kutetea haki za binadamu wa Bara la Afrika, Shirika la Amnesty International, Shirika la ARTICLE 19, Shirika la Asian Forum for Human Rights and Development, Kituo cha Centre for Civil Liberties – Ukraine, Human Rights Watch na Tume ya International Commission of Jurists na taasisi nyingine.
5142                                                                                                                                                                                                                                                                                                                                                           Rais Trump  Rais wa Marekani Donald Trump, anasema Biden, “amerusha baruti katika moto, na anajukumu la kutoa maelezo kwa watu wa Marekani juu ya mkakati na mpango wa kuhakikisha vikosi vyetu na wafanyakazi wa ubalozi, watu wetu na maslahi yetu, yote hapa nchini na nchi za nje, na washirika wetu katika eneo lote la Mashariki ya Kati na maeneo mengine.
5682                                                                                                                                                                                                                         ”  Ogwell amesema taasisi yake, ambayo ni shirika la ushauri wa kifundi la Umoja wa Afrika, anashirikiana na AU kuzijengea uwezo wa utayari nchi mbalimbali katika maeneo makuu matatu, ikiwemo kuboresha utoaji tahadhari katika bandari za nchi hizo na mahospitali; kuongeza utaalamu wa kuweza kupima kirusi COVID-19, ambao tayari nchi 43 wanauwezo huo; na kujenga uwezo wa kuzuia maambukizi na kudhibiti hali hiyo ili wagonjwa wenye maambukizi waweze kuwekewa karantini na kufuatiliwa.
4934                                                                                                                                                                                                                                                                                                                                                                             Morales alisema : "Kaka na dada zangu nchini Bolivia na ulimwenguni kote, nawafahamisha niko hapa na Makamu wa Rais na Waziri wa Afya, na baada ya kuwasilikiliza rafiki zangu kutoka shirikisho la vuguvugu la kijamii na shirikisho la umoja wa kibiashara na pia kwa kusikiliza Kanisa Katoliki, natangaza kujiuzulu wadhifa wangu wa urais.
7355                                                                                                                                                                                                                                                                                                         Kiongozi huyo wa cheo cha juu katika nchi za falme za kiarabu aliwaambia waandishi wa habari alisikiliza mtazamo wa Jenerali Abdel Fattah Burhan kuhusu matatizo ya Sudan na yeye alimweleza mtazamo wa umoja wa falme za kiarabu kuhusiana na hali hii ya kisiasa nchini Sudan Katibu mkuu wa umoja wa nchi za falme za kiarabu alifanya mazungumzo mjini Khartoum Jumapili na baraza la jeshi linalotawala Sudan.
5437                                                                                                                                                                                                                                                                                                                                                                                                                                                                                                             Pamoja na kuwa wanashirikiana katika mipaka yao na kuwepo kwao katika wigo la kiuchumi la pamoja, nchi za Afrika Mashariki zinakuwa na hisia kali baina yao zinazotokana na tofauti zao za kiuchumi na kisiasa.
</pre>

 <p>
The results don't look too relevant, probably because they just happen
to be long sentences that contain some words from the query several
times.  The dot product is sensitive to vector magnitudes, so longer
sentences are likely to score higher just because they have a higher
count of words in the query.
</p>

 <p>
 <a href="https://web.archive.org/web/20191213082655/https://www.sciencedirect.com/topics/computer-science/cosine-similarity">Cosine similarity</a> normalizes for magnitude of the vectors making the
score less sensitive to the absolute counts of similar terms.
</p>

 <p>
Results for cosine similarity:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.metrics.pairwise  <span style="font-weight: bold;">import</span> cosine_similarity

 <span style="font-weight: bold; font-style: italic;">score</span> = cosine_similarity(X, q).flatten()
df.iloc[np.argsort(-score)[:10]]
</pre>
</div>

 <pre class="example" id="org9f04347">
                                                                                                                                                                                                                                                                                                                                                                                   documents
5437                                                                                                                                                                         Pamoja na kuwa wanashirikiana katika mipaka yao na kuwepo kwao katika wigo la kiuchumi la pamoja, nchi za Afrika Mashariki zinakuwa na hisia kali baina yao zinazotokana na tofauti zao za kiuchumi na kisiasa.
3022                                                                                                                                                                                                                   “Inasikitisha kwamba nchi kadhaa duniani na mashirika mbalimbali, ikiwemo Marekani, China na shirika la afya duniani, yanakabiliwa na tishio la maambukizi ya Corona.
4808                                                                                                                                                               Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
664                                                                                                                                                                Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
6369                                                                                                                 Mkurugenzi Mkuu wa Shirika la Afya Duniani Tedros Adhanom Ghebreyesus, ameonya kufungwa kwa mipaka ya nchi na kusitisha shughuli zote ili kupambana na janga la COVID 19 kunaweza kusababisha kuongezeka kwa vifo kutokana na ugonjwa wa Malaria katika nchi za Afrika.
4571                                                                                                                                                                                                           Shahidi mmoja aliliambia shirika la habari la Reuters katika wiki kadhaa za karibuni za maandamano yanayoipinga serikali yaliyochochewa na malalamiko ya kiuchumi na kisiasa.
6182                                                                                                                    Ndege za kijeshi za India, zilivuka mpaka na kuingia katika nchi jirani ya Pakisan na kutekeleza mashambulizi dhidi ya kambi iliyodaiwa kuwa ya kutoa mafunzi kwa kundi la wanamgambo la Jaish-e-Mohammad, lililoripotiwa kuhusika na shambulizi la bomu la Kashmir.
6107  Lissu ameongeza kwamba “Kwa kutumia njia hizo za amri za kisiasa, vyombo vya ulinzi na usalama vikiwemo jeshi la wananchi la Tanzania, jeshi la polisi, taasisi ya kuzuia na kupambana na rushwa, na idara ya usalama wa taifa pamoja na mamlaka ya kodi TRA, vinatumika kukamata kwa nguvu na kutaifisha mali za wafanyabiashara wetu wa ndani na makampuni ya wawekezaji kutoka nje.
3191                                                                                                                                SAA, shirika kubwa la ndege la Afrika liliingia katika mpango wa kujilinda kutokana na hali ya kufilisika mwezi Disemba 2019, na tangu wakati huo lililazimika kusitisha safari zake zote za abiria kutokana na janga la virusi vya korona kote duniani.
4993                                                                                                                                                                                                                                                                     Shirika la habari la China Xinhua linasema kutakuwa na karibu safari 200 za ndege kuingia na kutoka Wuhan Jumatano.
</pre>

 <p>
Cosine similarity results in more relevant results.
</p>

 <p>
We create a generic function for getting search results from the
vector space like this:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">search</span>(X, query, num_results=10):
     <span style="font-weight: bold; font-style: italic;">score</span> = cosine_similarity(X, query).flatten()
     <span style="font-weight: bold; font-style: italic;">idx</span> = np.argsort(-score)[:num_results]
     <span style="font-weight: bold;">return</span> df.iloc[idx]
</pre>
</div>
</div>
</li>
</ol></div>
 <div id="outline-container-org6255a02" class="outline-4">
 <h4 id="org6255a02"> <span class="section-number-4">4.2.2.</span> TF-IDF Vectorization</h4>
 <div class="outline-text-4" id="text-4-2-2">
 <p>
There are a number of common words like: na, kwa, la, ya, etc, which
are common across the entire corpus but are relatively insignificant
to the relevance of a particular document to the query.
</p>

 <p>
TF-IDF or  <a href="https://en.wikipedia.org/wiki/Tf%E2%80%93idf">term frequency-inverse document frequency</a> in full, minimises
the effect of these words. It introduces a new score in place of
counts in the document-term matrix that shows how important a term is
to the document.
</p>

 <p>
The score is calculated by multiplying the  <b>term frequency</b> - the
frequency of the term in relation to other terms in the document, and
the  <b>inverse document frequency</b> - how rare the term is across all
the documents in the corpus.
</p>

 <p>
\[
\text{TF}(t, d) = \frac{\text{count of } t \text{ in } d}{\text{total terms in } d}
\]
</p>

 <p>
\[
\text{IDF}(t) = \log\frac{\text{total documents}}{\text{number of documents containing } t}
\]
</p>

 <p>
To do this, we replace the Count Vectorizer with a TfidVectorizer:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.feature_extraction.text  <span style="font-weight: bold;">import</span> TfidfVectorizer

 <span style="font-weight: bold; font-style: italic;">tfv</span> = TfidfVectorizer()
 <span style="font-weight: bold; font-style: italic;">X_tfv</span> = tfv.fit_transform(df.documents)
 <span style="font-weight: bold; font-style: italic;">q_tfv</span> = tfv.transform([query])
 <span style="font-weight: bold; font-style: italic;">names</span> = tfv.get_feature_names_out()
</pre>
</div>

 <p>
Display the non-zero terms in the first document:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">df_docs</span> = pd.DataFrame(X_tfv.toarray(), columns=names)
 <span style="font-weight: bold; font-style: italic;">first_doc</span> = df_docs.loc[0]
first_doc[first_doc != 0]
</pre>
</div>

 <pre class="example" id="org7321461">
14            0.311283
19            0.260709
afya          0.234890
covid         0.273393
imeripoti     0.327010
jumatatu      0.253012
kuwa          0.143414
maambukizi    0.222536
takriban      0.275575
tanzania      0.210560
wamepata      0.413453
watu          0.161859
wizara        0.255201
ya            0.216513
zaidi         0.186419
Name: 0, dtype: float64
</pre>

 <p>
We no longer have counts, but the tf-idf score.
</p>

 <p>
We can now do a TF-IDF search which should show improved results:
</p>

 <div class="org-src-container">
 <pre class="src src-python">search(X_tfv, q_tfv)
</pre>
</div>

 <pre class="example" id="org20a21ed">
                                                                                                                                                                                                                                                                    documents
664                                                 Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
4808                                                Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
2574                                                                                                                                                                  Mali, nchi yenye huduma mbaya za afya, iliandaa uchaguzi jumapili, licha ya janga la virusi vya Corona.
6369  Mkurugenzi Mkuu wa Shirika la Afya Duniani Tedros Adhanom Ghebreyesus, ameonya kufungwa kwa mipaka ya nchi na kusitisha shughuli zote ili kupambana na janga la COVID 19 kunaweza kusababisha kuongezeka kwa vifo kutokana na ugonjwa wa Malaria katika nchi za Afrika.
7116                                           China imesema Jumatano kuwa uamuzi wa Rais wa Marekani Donald Trump kusitisha ufadhili kwa Shirika la Afya Duniani kutaziathiri nchi zote wakati dunia ikikabiliwa na hatua muhimu ya kupambana na janga la virusi vya corona.
2311                                                                                           Hatua hiyo imechukuliwa huku maambukizi yakiendelea kuongezeka kote duniani, na baada ya Shirika la Afya Duniani (WHO) kutangaza maambukizi ya Corona kuwa janga la kimataifa.
3022                                                                                                    “Inasikitisha kwamba nchi kadhaa duniani na mashirika mbalimbali, ikiwemo Marekani, China na shirika la afya duniani, yanakabiliwa na tishio la maambukizi ya Corona.
814                                              Katika ukosoaji wa nadra kwa umma, shirika la afya Duniani-WHO wiki iliyopita lilieleza kwamba katika kupingana na kanuni za kimataifa za afya, Tanzania ilikataa kutoa taarifa za kina za kesi zinazoshukiwa kuwa za Ebola.
911                                                                                                                              Wakati huo huo, Museveni ametangaza mipango ya kuwarudisha Uganda raia wa nchi hiyo ambao wamekwama nchi za nje kutokana na janga la Corona.
1020                                                                                                                                            Kwa mujibu wa shirika la habari la Uingereza Reuters hisa za shirika la ndege la Ujerumani Luftansa zilipanda kwa asilimia 6.
</pre>
</div>
</div>
</div>
 <div id="outline-container-orgb3cd5e4" class="outline-3">
 <h3 id="orgb3cd5e4"> <span class="section-number-3">4.3.</span> Embeddings</h3>
 <div class="outline-text-3" id="text-4-3">
 <p>
A problem we still have is that we're matching for exact terms in the
documents i.e. lexical search, therefore synonyms and closely related
terms won't be captured in the search.
</p>

 <p>
To fix this we use embeddings, which cluster related words
together to capture ideas or concepts, and contextual information.
</p>
</div>
 <div id="outline-container-org914ce40" class="outline-4">
 <h4 id="org914ce40"> <span class="section-number-4">4.3.1.</span> Singular Value Decomposition (SVD)</h4>
 <div class="outline-text-4" id="text-4-3-1">
 <p>
This is a technique in linear algebra that operates on a matrix to extract
its most important features. It can be used, for example, in  <a href="https://timbaumann.info/svd-image-compression-demo/">lossy
compression of images</a> where important features of the image are
extracted, and can be used to recreate the original image but with a lower
resolution.
</p>

 <p>
When used on a document-term matrix, it extracts association between
related words that represent a concept or topic.
</p>

 <p>
It reduces the dimensionality and captures the underlying semantic
structure by grouping together related words, effectively extracting
concepts or topics expressed by co-occurrence patterns in the corpus.
</p>

 <p>
We use the vector representation from the TF-IDF vectorizer to create
SVD embeddings:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.decomposition  <span style="font-weight: bold;">import</span> TruncatedSVD

 <span style="font-weight: bold; font-style: italic;">svd</span> = TruncatedSVD(n_components=500)
 <span style="font-weight: bold; font-style: italic;">X_svd</span> = svd.fit_transform(X_tfv)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">X_svd.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">7672</td>
 <td class="org-right">500</td>
</tr></tbody></table> <p>
We provided the number of topics to be extracted as 500. For each
document, the matrix shows how much it ranks in importance to each of
the 500 topics.
</p>

 <p>
We can see how the first document ranks for the first 5 topics:
</p>

 <div class="org-src-container">
 <pre class="src src-python">X_svd[0,:5]
</pre>
</div>

 <pre class="example">
array([ 0.18967023, -0.06791604, -0.24394785, -0.04354468,  0.00576036])
</pre>


 <p>
The  <code>n_components</code> parameter is dependent on the dataset used and some
experimentation may be required to arrive at the optimal value.
</p>

 <p>
We similarly create embeddings for the query and run a search:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">q_svd</span> = svd.transform(q_tfv)
q_svd.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">1</td>
 <td class="org-right">500</td>
</tr></tbody></table> <div class="org-src-container">
 <pre class="src src-python">search(X_svd, q_svd)
</pre>
</div>

 <pre class="example" id="org1574667">
                                                                                                                                                                                                                                                                                                                                                                                                                                                                                          documents
664                                                                                                                                                                                                                                                                       Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
4808                                                                                                                                                                                                                                                                      Amesema "mwanzoni tulizungumza kujifunza kutoka uzoefu wa nchi nyingine katika kupambana na janga la corona, sasa tunazungumzwa na nchi nyingine kama kielelezo hasi cha vita dhidi ya janga la corona Afrika na duniani.
6369                                                                                                                                                                                                                        Mkurugenzi Mkuu wa Shirika la Afya Duniani Tedros Adhanom Ghebreyesus, ameonya kufungwa kwa mipaka ya nchi na kusitisha shughuli zote ili kupambana na janga la COVID 19 kunaweza kusababisha kuongezeka kwa vifo kutokana na ugonjwa wa Malaria katika nchi za Afrika.
2574                                                                                                                                                                                                                                                                                                                                                                                        Mali, nchi yenye huduma mbaya za afya, iliandaa uchaguzi jumapili, licha ya janga la virusi vya Corona.
968   Akizungumza na shirika la habari la AFP Akram Taher mumini moja amesema “Sherehe za Eid hazifani manmo wakati huu wa hali ya janga la corona – watu wanahisia ya kua na hofu”  Bara la Asia  Waislam katika bara la Asia – kutoka Indonesia hadi Pakistan, Malaysia na Afghanistan – wamekusanyika katika masoko katika shamrashamra za manunuzi ya sikukuu, wakikiuka muongozo wa kudhibiti virusi vya corona na wakati mwengine polisi wakijaribu kutawanya mikusanyiko ya makundi makubwa.
814                                                                                                                                                                                                                                                                    Katika ukosoaji wa nadra kwa umma, shirika la afya Duniani-WHO wiki iliyopita lilieleza kwamba katika kupingana na kanuni za kimataifa za afya, Tanzania ilikataa kutoa taarifa za kina za kesi zinazoshukiwa kuwa za Ebola.
3022                                                                                                                                                                                                                                                                                                                          “Inasikitisha kwamba nchi kadhaa duniani na mashirika mbalimbali, ikiwemo Marekani, China na shirika la afya duniani, yanakabiliwa na tishio la maambukizi ya Corona.
7116                                                                                                                                                                                                                                                                 China imesema Jumatano kuwa uamuzi wa Rais wa Marekani Donald Trump kusitisha ufadhili kwa Shirika la Afya Duniani kutaziathiri nchi zote wakati dunia ikikabiliwa na hatua muhimu ya kupambana na janga la virusi vya corona.
1505                                                                                                                                                                                                                                                                                                                              Nchi hizo za Afrika magharibi zimekuwa zikirekodi ongezeko la watu wanaoambukizwwa virusi vya Corona, na haijulikani namna zitakavyokabiliana na maambukizi hayo.
4095                                                                                                                                                                                                                                                                                            Wakati huo huo, Umoja wa Mataifa (UN) umeonya kwamba janga la virusi vya corona linaendelea kusababisha matatizo ya kiakili na msongo wa mawazo, hasa katika nchi ambazo kuna sekta dhaifu za afya.
</pre>
</div>
</div>
 <div id="outline-container-org2166ff3" class="outline-4">
 <h4 id="org2166ff3"> <span class="section-number-4">4.3.2.</span> BERT</h4>
 <div class="outline-text-4" id="text-4-3-2">
 <p>
The previous methods all used the bag-of-words approach; the order of
words wasn't taken into consideration when searching. The order of the
words in the documents may add contextual useful information that
could improve the search.
</p>

 <p>
BERT is a deep neural network model of the transformer architecture
that encodes contextual meaning of words taking into account where
they occur in a sentence, having being pre-trained on a large corpus
of text.
</p>

 <p>
We'll use a variant of BERT called
 <code>flax-community/bert-swahili-news-classification</code>, which has been
fine-tuned on Swahili news text, that will be downloaded from  <a href="https://huggingface.co/flax-community/bert-swahili-news-classification">Hugging
Face</a>.
</p>

 <p>
Each pretrained transformer model has a tokenizer that is used to encode
the input text into the vocabulary that was used in the training. We
instantiate both the tokenizer and the model:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">Load model directly
</span> <span style="font-weight: bold;">from</span> transformers  <span style="font-weight: bold;">import</span> AutoTokenizer, AutoModelForSequenceClassification

 <span style="font-weight: bold; font-style: italic;">tokenizer</span> = AutoTokenizer.from_pretrained( <span style="font-style: italic;">"flax-community/bert-swahili-news-classification"</span>)
 <span style="font-weight: bold; font-style: italic;">model</span> = AutoModelForSequenceClassification.from_pretrained( <span style="font-style: italic;">"flax-community/bert-swahili-news-classification"</span>)
</pre>
</div>

 <p>
Using two documents from our text corpus, we run through the process
of creating BERT embeddings:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">texts</span> = df.documents[:2].tolist()
 <span style="font-weight: bold; font-style: italic;">encoded_text</span> = tokenizer(texts, padding= <span style="font-weight: bold; text-decoration: underline;">True</span>, return_tensors= <span style="font-style: italic;">'pt'</span>)
</pre>
</div>

 <p>
We can see the encoded documents as the  <code>input_ids</code> attribute:
</p>

 <div class="org-src-container">
 <pre class="src src-python">encoded_text.input_ids.shape
</pre>
</div>

 <pre class="example">
torch.Size([2, 23])
</pre>



 <div class="org-src-container">
 <pre class="src src-python">encoded_text.input_ids
</pre>
</div>

 <pre class="example">
tensor([[    2,  1057,   117,   902,   117,   367,   362,  2867,  3731,   200,
            20,   283,  5271,   869,   349,  8119,  4252,   117,  7632,    21,
           588,    22,     3],
        [    2, 15117,  7672,   587,   156,  1938,   115,   367,    20,   870,
          1374,   544,    21,   595,    21,   662,   119,   628,   969,  1617,
            22,     3,     0]])
</pre>


 <p>
The document embeddings will be found at the last hidden layer before the
output layer, after doing a forward pass through the model:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">import</span> torch

model.config. <span style="font-weight: bold; font-style: italic;">output_hidden_states</span>= <span style="font-weight: bold; text-decoration: underline;">True</span>
 <span style="font-weight: bold;">with</span> torch.no_grad():   <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">Disable gradient calculation since we aren't training
</span>     <span style="font-weight: bold; font-style: italic;">outputs</span> = model(**encoded_text)
     <span style="font-weight: bold; font-style: italic;">last_hidden_states</span> = outputs.hidden_states[-1]
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">last_hidden_states.shape
</pre>
</div>

 <pre class="example">
torch.Size([2, 23, 768])
</pre>


 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">compressed_emb</span> = last_hidden_states.mean(dim=1)
compressed_emb.shape
</pre>
</div>

 <pre class="example">
torch.Size([2, 768])
</pre>



 <div class="org-src-container">
 <pre class="src src-python">compressed_emb.numpy()
</pre>
</div>

 <pre class="example">
array([[ 1.4375398 , -0.24863385, -0.06177869, ..., -1.2794679 ,
        -0.58472764,  0.7913138 ],
       [ 0.536299  , -0.8806017 , -0.8250647 , ..., -0.3449616 ,
         0.45784512,  0.97287774]], shape=(2, 768), dtype=float32)
</pre>


 <p>
At this point, we have a representation of 768 topic scores for each
input document.
</p>

 <p>
To repeat the process for the entire set of documents, we can do
inference in batches so as not to overwhelm the hardware. Modify
 <code>batch_size</code> as appropriate for your machine.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> tqdm  <span style="font-weight: bold;">import</span> tqdm

 <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">get_embeddings</span>(documents, batch_size=50):
     <span style="font-weight: bold; font-style: italic;">embedding_batches</span> = []
     <span style="font-weight: bold;">for</span> i  <span style="font-weight: bold;">in</span> tqdm( <span style="font-weight: bold;">range</span>(0,  <span style="font-weight: bold;">len</span>(documents), batch_size)):
         <span style="font-weight: bold; font-style: italic;">batch</span> = documents[i:i+batch_size]
         <span style="font-weight: bold; font-style: italic;">encoded_text</span> = tokenizer(batch, padding= <span style="font-weight: bold; text-decoration: underline;">True</span>, return_tensors= <span style="font-style: italic;">'pt'</span>)
         <span style="font-weight: bold;">with</span> torch.no_grad():
             <span style="font-weight: bold; font-style: italic;">outputs</span> = model(**encoded_text)
             <span style="font-weight: bold; font-style: italic;">last_hidden_states</span> = outputs.hidden_states[-1]
            embedding_batches.append(last_hidden_states.mean(dim=1).numpy())
     <span style="font-weight: bold;">return</span> np.vstack(embedding_batches)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">X_bert</span> = get_embeddings(df.documents.to_list())
</pre>
</div>

 <p>
We also convert the query into embeddings:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">q_bert</span> = get_embeddings([query], batch_size=1)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">X_bert.shape, q_bert.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">7672</td>
 <td class="org-right">768</td>
</tr> <tr> <td class="org-right">1</td>
 <td class="org-right">768</td>
</tr></tbody></table> <p>
Then we run a search using BERT embeddings:
</p>

 <div class="org-src-container">
 <pre class="src src-python">search(X_bert, q_bert)
</pre>
</div>

 <pre class="example" id="orgcc7af20">
                                                                                                                                                                                                                                              documents
3022                                                                              “Inasikitisha kwamba nchi kadhaa duniani na mashirika mbalimbali, ikiwemo Marekani, China na shirika la afya duniani, yanakabiliwa na tishio la maambukizi ya Corona.
5679                                     Nchi tatu za Africa zinamaambukizi zaidi katika mlipuko wa virusi vya corona bara la Afrika, lakini naibu mkurugenzi wa Vituo vya Kudhibiti Magonjwa na Kinga Afrika wanasema bara lote lazima lichukue hatua.
4832        Janga la virusi vya corona limeleta ubunifu kwa watu wengi duniani ambao nchi zao zimetoa amri ya kutotoka majumbani, na vilevile wanaendelea kutoa misaada kwa wafanyakazi wa afya ambao wanawatibu wale waliopata maambukizi ya COVID-19.
6025                                                   Idadi ya maambukizi ya virusi vya corona yamegunduliwa katika nchi zaidi wakati jjumiya ya kimataifa na serikali mbalimbali zikijaribu kuhakikisha wanadhibiti maambukizi mapya katika nchi zao.
1481                                                                                             Wizara ya afya a Uganda imesema kwamba shirika la afya duniani limetoa mwelekeo kwamba kila kisa cha maambukizi kinahesabiwa katika nchi kimeripotiwa.
3016                                                                            Baadhi ya viongozi barani Afrika wameeleza kushangazwa kwao na tamko la rais wa Marekani kwamba anafikiria kusitisha msaada wa kifedha kwa Shirika la Afya Duniani WHO.
4090                                                                                                                                  Shirika la Afya Duniani, WHO, limetahadharisha uwezekano wa kuendelea virusi vya corona vikaendelea kuwepo daima.
2460                                                            Shirika la Afya Duniani (WHO) Ijumaa limeonya kuwa watu 190,000 wanaweza kupoteza maisha mwaka 2020 barani Africa, iwapo serikali zitashindwa kudhibiti maambukizi ya virus vya corona.
6932                                                           Mkuu wa Shirika la Afya, WHO, nchini Burundi alifukuzwa wiki iliyopita baada ya kueleza wasiwasi wake juu ya uelewa wa serikali unavyokinzana na hatari inayoletwa na virusi vya corona.
2541  “Wizara ya afya, shirika la afya duniano na kituo cha kukabiliana na magonjwa ya kuambikaza, wataanza kutoa chanjo dhidi ya ebola kwa watu wanaoaminika kuwa karibu na visa vilivyothibitishwa kabla ya chanjo kutolewa kwa uma, kuanzia juni 14.
</pre>

 <p>
The BERT results seem to be the most relevant so far.
</p>
</div>
</div>
</div>
</div>
 <div id="outline-container-org00b0096" class="outline-2">
 <h2 id="org00b0096"> <span class="section-number-2">5.</span> Conclusion</h2>
 <div class="outline-text-2" id="text-5">
 <p>
We looked at various techniques that can be used to search textual
documents, starting from a simple keyword based approach to a more
sophisticated one utilizing language model embeddings like BERT in an
effort to improve the results. These are a subset of techniques that
are applied in the wider field of Information Retrieval. A good
resource for a more in-depth introduction to the fields is found on
the  <a href="https://nlp.stanford.edu/IR-book/information-retrieval-book.html">Stanford NLP website</a>.
</p>
</div>
</div>
</div>]]></content>
  <link href="https://dataai.ng/search_swahili.html"/>
  <id>https://dataai.ng/search_swahili.html</id>
  <updated>2026-07-31T15:07:00+03:00</updated>
</entry>
<entry>
  <title>Creating a Digit Classifier (Almost) from Scratch</title>
  <author><name>KE Programmer</name></author>
  <summary>Published on Feb 21, 2023 12:36 by KE Programmer</summary>
  <content type="html"><![CDATA[<div id="content" class="content">
 <h1 class="title">Creating a Digit Classifier (Almost) from Scratch
 <br></br> <span class="subtitle">Published on Feb 21, 2023 12:36 by KE Programmer</span>
</h1>
 <div id="table-of-contents" role="doc-toc">
 <h2>Table of Contents</h2>
 <div id="text-table-of-contents" role="doc-toc">
 <ul> <li> <a href="#org1e47510">1. Introduction</a></li>
 <li> <a href="#org99b2320">2. Data Prep</a></li>
 <li> <a href="#orge5a0c96">3. Training</a>
 <ul> <li> <a href="#orgff4f698">3.1. Pytorch/FastAI conveniences</a></li>
</ul></li>
 <li> <a href="#org8ca759e">4. Conclusion</a></li>
</ul></div>
</div>
 <div id="outline-container-org1e47510" class="outline-2">
 <h2 id="org1e47510"> <span class="section-number-2">1.</span> Introduction</h2>
 <div class="outline-text-2" id="text-1">
 <p>
We're creating a model that can classify any images as a 3 or
a 7. We'll use a sample of MNIST that contains just these.
</p>
</div>
</div>
 <div id="outline-container-org99b2320" class="outline-2">
 <h2 id="org99b2320"> <span class="section-number-2">2.</span> Data Prep</h2>
 <div class="outline-text-2" id="text-2">
 <p>
First, let's import the necessary libraries:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> fastai.vision. <span style="font-weight: bold;">all</span>  <span style="font-weight: bold;">import</span> *
 <span style="font-weight: bold;">import</span> uuid
 <span style="font-weight: bold;">import</span> os
 <span style="font-weight: bold;">import</span> pathlib
</pre>
</div>

 <p>
Download and explore the data:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">path</span> = untar_data(URLs.MNIST_SAMPLE)
path.ls()
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/labels.csv)</td>
 <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/train)</td>
 <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/valid)</td>
</tr></tbody></table> <p>
The sample data is divided into training and validation sets:
</p>

 <div class="org-src-container">
 <pre class="src src-python">(path/ <span style="font-style: italic;">"train"</span>).ls(), (path/ <span style="font-style: italic;">"valid"</span>).ls()
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/train/7)</td>
 <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/train/3)</td>
</tr> <tr> <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/valid/7)</td>
 <td class="org-left">Path</td>
 <td class="org-left">( <i>home/krm</i>.fastai/data/mnist <sub>sample</sub>/valid/3)</td>
</tr></tbody></table> <p>
Get a list of the training set of 3s and 7s:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">threes</span> = (path/ <span style="font-style: italic;">"train"</span>/ <span style="font-style: italic;">"3"</span>).ls(). <span style="font-weight: bold;">sorted</span>()
 <span style="font-weight: bold; font-style: italic;">sevens</span> = (path/ <span style="font-style: italic;">"train"</span>/ <span style="font-style: italic;">"7"</span>).ls(). <span style="font-weight: bold;">sorted</span>()
</pre>
</div>

 <p>
Have a look at one of the 3s. We use the  <code>Image</code> class from PIL:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">img3_path</span> = threes[2]
 <span style="font-weight: bold; font-style: italic;">img3</span> = Image. <span style="font-weight: bold;">open</span>(img3_path)
img3
</pre>
</div>



 <div id="org56b38a8" class="figure">
 <p> <img src="img/mnist_3_sample.png" alt="mnist_3_sample.png"></img></p>
</div>

 <p>
We can see the image as a collection of digits using numpy arrays or
pytorch tensors:
</p>

 <div class="org-src-container">
 <pre class="src src-python">array(img3)[4:10, 4:10]
</pre>
</div>

 <pre class="example">
array([[  0,   0,   0,   0,   0,   0],
       [  0,   0,   0,   0,   0,   0],
       [  0,   0,   0,   0,  13,  36],
       [  0,   0,   0,   0,  89, 253],
       [  0,   0,   0,   0,  89, 253],
       [  0,   0,   0,   0,  17, 151]], dtype=uint8)
</pre>


 <div class="org-src-container">
 <pre class="src src-python">tensor(img3)[4:10, 4:10]
</pre>
</div>

 <pre class="example">
tensor([[  0,   0,   0,   0,   0,   0],
        [  0,   0,   0,   0,   0,   0],
        [  0,   0,   0,   0,  13,  36],
        [  0,   0,   0,   0,  89, 253],
        [  0,   0,   0,   0,  89, 253],
        [  0,   0,   0,   0,  17, 151]], dtype=torch.uint8)
</pre>


 <p>
Next we create tensors for each of the 3s and 7s:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">three_tensors</span> = [tensor(Image. <span style="font-weight: bold;">open</span>(path))  <span style="font-weight: bold;">for</span> path  <span style="font-weight: bold;">in</span> threes]
 <span style="font-weight: bold; font-style: italic;">seven_tensors</span> = [tensor(Image. <span style="font-weight: bold;">open</span>(path))  <span style="font-weight: bold;">for</span> path  <span style="font-weight: bold;">in</span> sevens]
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">len</span>(three_tensors),  <span style="font-weight: bold;">len</span>(seven_tensors)
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">6131</td>
 <td class="org-right">6265</td>
</tr></tbody></table> <p>
Use fastai's  <code>show_image</code> to see a seven:
</p>

 <div class="org-src-container">
 <pre class="src src-python">show_image(seven_tensors[6])
</pre>
</div>


 <div id="org5589f60" class="figure">
 <p> <img src="img/mnist_7_sample.png" alt="mnist_7_sample.png"></img></p>
</div>

 <p>
We stack each list of tensors into a single one of 3 axes (rank 3),
convert to float for some operations. We also scale to values between
0 and 1 (better for the model):
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">stacked_threes</span> = torch.stack(three_tensors). <span style="font-weight: bold;">float</span>() / 255
 <span style="font-weight: bold; font-style: italic;">stacked_sevens</span> = torch.stack(seven_tensors). <span style="font-weight: bold;">float</span>() / 255
stacked_threes.shape, stacked_sevens.shape
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">torch.Size</td>
 <td class="org-left">((6131 28 28))</td>
 <td class="org-left">torch.Size</td>
 <td class="org-left">((6265 28 28))</td>
</tr></tbody></table> <p>
We create a training input collection by concatenating the two stacked
tensors into one. We use the  <code>view</code> method to reshape the tensor into
two dimensions.  <code>-1</code> means making the first dimension as big as
possible to accommodate the new shape:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">train_x</span> = torch.cat([stacked_threes, stacked_sevens]).view(-1, 28 * 28)
train_x.shape
</pre>
</div>

 <pre class="example">
torch.Size([12396, 784])
</pre>


 <div class="org-src-container">
 <pre class="src src-python">train_x[:2]
</pre>
</div>

 <pre class="example">
tensor([[0., 0., 0.,  ..., 0., 0., 0.],
        [0., 0., 0.,  ..., 0., 0., 0.]])
</pre>


 <p>
Each item along axis 0 is now a list of 784 floats representing a
single image.
</p>

 <p>
Next we create labels for our training input. Our objective is to
classify whether an image is a 3 or not. The labels for 3s will be 1
(True) and for 7s will be 0 (False):
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">train_y_flat</span> = tensor([1] *  <span style="font-weight: bold;">len</span>(threes) + [0] *  <span style="font-weight: bold;">len</span>(sevens))
train_y_flat.shape
</pre>
</div>

 <pre class="example">
torch.Size([12396])
</pre>


 <p>
We need to have a 2D tensor with the second dimension being a size
of 1. We use pytorch's  <code>unsqueeze</code> for that:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">train_y</span> = train_y_flat.unsqueeze(1)
train_y.shape
</pre>
</div>

 <pre class="example">
torch.Size([12396, 1])
</pre>


 <p>
In pytorch, a dataset needs to return a tuple of (x, y) when
indexed. We therefore zip training input and labels:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">dataset</span> =  <span style="font-weight: bold;">list</span>( <span style="font-weight: bold;">zip</span>(train_x, train_y))
 <span style="font-weight: bold; font-style: italic;">x0</span>,  <span style="font-weight: bold; font-style: italic;">y0</span> = dataset[0]
x0.shape, y0.shape
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">torch.Size</td>
 <td class="org-left">((784))</td>
 <td class="org-left">torch.Size</td>
 <td class="org-left">((1))</td>
</tr></tbody></table> <p>
We do the same preparation to the validation data as we've done for
the training data:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">v_threes</span> = (path/ <span style="font-style: italic;">"valid"</span>/ <span style="font-style: italic;">"3"</span>).ls(). <span style="font-weight: bold;">sorted</span>()
 <span style="font-weight: bold; font-style: italic;">v_sevens</span> = (path/ <span style="font-style: italic;">"valid"</span>/ <span style="font-style: italic;">"7"</span>).ls(). <span style="font-weight: bold;">sorted</span>()
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">v_three_tensors</span> = [tensor(Image. <span style="font-weight: bold;">open</span>(path))  <span style="font-weight: bold;">for</span> path  <span style="font-weight: bold;">in</span> v_threes]
 <span style="font-weight: bold; font-style: italic;">v_seven_tensors</span> = [tensor(Image. <span style="font-weight: bold;">open</span>(path))  <span style="font-weight: bold;">for</span> path  <span style="font-weight: bold;">in</span> v_sevens]
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">stacked_v_threes</span> = torch.stack(v_three_tensors). <span style="font-weight: bold;">float</span>()/255
 <span style="font-weight: bold; font-style: italic;">stacked_v_sevens</span> = torch.stack(v_seven_tensors). <span style="font-weight: bold;">float</span>()/255
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">valid_x</span> = torch.cat([stacked_v_threes, stacked_v_sevens]).view(-1, 28 * 28)
valid_x.shape
</pre>
</div>

 <pre class="example">
torch.Size([2038, 784])
</pre>


 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">valid_y</span> = tensor([1] *  <span style="font-weight: bold;">len</span>(v_threes) + [0] *  <span style="font-weight: bold;">len</span>(v_sevens)).unsqueeze(1)
valid_y.shape
</pre>
</div>

 <pre class="example">
torch.Size([2038, 1])
</pre>


 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">validation_dataset</span> =  <span style="font-weight: bold;">list</span>( <span style="font-weight: bold;">zip</span>(valid_x, valid_y))
</pre>
</div>
</div>
</div>
 <div id="outline-container-orge5a0c96" class="outline-2">
 <h2 id="orge5a0c96"> <span class="section-number-2">3.</span> Training</h2>
 <div class="outline-text-2" id="text-3">
 <p>
The initial model will be a linear function with weights for each
pixel and a bias. We'll need to calculate gradients for each of these,
so we call  <code>requires_grad</code>:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">init_params</span>(size, std=1.0):  <span style="font-weight: bold;">return</span> (torch.randn(size) * std).requires_grad_()
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">weights</span> = init_params((28*28, 1))
 <span style="font-weight: bold; font-style: italic;">bias</span> = init_params(1)
weights.shape, bias.shape
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">torch.Size</td>
 <td class="org-left">((784 1))</td>
 <td class="org-left">torch.Size</td>
 <td class="org-left">((1))</td>
</tr></tbody></table> <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">linear</span>(x_batch):  <span style="font-weight: bold;">return</span> x_batch@weights + bias
</pre>
</div>

 <p>
The  <code>@</code> operator does the dot-product operation  
</p>

 <pre class="example">
None
</pre>


 <p>
Next, we need to define a loss function that is sensitive to small
changes in the parameters. For each batch of predictions, we calculate
distance from the target values, and get the mean.
</p>

 <p>
We also need to coerce the predictions to values between 0 and 1, so
we'll use a sigmoid function for this:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">sigmoid</span>(x):  <span style="font-weight: bold;">return</span> 1/(1+torch.exp(-x))
</pre>
</div>

 <p>
How it does this is that for large positive values of  <code>x</code>,
 <code>torch.exp(-x)</code> approaches zero and the function outputs a value
approaching 1. For large negative values of  <code>x</code>,  <code>torch.exp(-x)</code>
output a large positive number the the output approaches 0.
</p>

 <div class="org-src-container">
 <pre class="src src-python">sigmoid(tensor(34343424)), sigmoid(tensor(-385984549058))
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">tensor</td>
 <td class="org-left">(1)</td>
 <td class="org-left">tensor</td>
 <td class="org-left">(0)</td>
</tr></tbody></table> <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">mnist_loss</span>(predictions, targets):
     <span style="font-weight: bold; font-style: italic;">predictions</span> = sigmoid(predictions)
     <span style="font-weight: bold;">return</span> torch.where(targets==1, 1-predictions, predictions).mean()
</pre>
</div>

 <p>
The value returned by  <code>minst_loss</code> is bounded in [0, 1], so the closer
it is to 1, the worse the prediction is.
</p>

 <p>
We now have enough to calculate gradients. We define a procedure that
makes predictions, calculates the loss, then calculates the gradients
that would minimise the loss:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">calculate_gradients</span>(x_batch, y_batch, model):
     <span style="font-weight: bold; font-style: italic;">predictions</span> = model(x_batch)
     <span style="font-weight: bold; font-style: italic;">loss</span> = mnist_loss(predictions, y_batch)
    loss.backward()
</pre>
</div>

 <p>
We need to iteratively process mini batches until we exhaust the
entire dataset. The  <code>fastai</code> library provides a  <code>DataLoader</code> class
that will shuffle the dataset and provide batches of input and their
corresponding targets according to the batch size we provide. We can
iterate through this to process the entire dataset:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">train_dl</span> = DataLoader(dataset, batch_size=256)
 <span style="font-weight: bold; font-style: italic;">valid_dl</span> = DataLoader(validation_dataset, batch_size=256)
</pre>
</div>

 <p>
Now we can encode the process of training through an entire
epoch. With each mini-batch, we optimize the parameters using the
gradients and learning rate. We're adjusting each parameter in
opposite direction of the gradient to get closer to a minimized loss:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">train_epoch</span>(model, learning_rate, params):
     <span style="font-weight: bold;">for</span> x, y  <span style="font-weight: bold;">in</span> train_dl:
        calculate_gradients(x, y, model)
         <span style="font-weight: bold;">for</span> p  <span style="font-weight: bold;">in</span> params:
            p. <span style="font-weight: bold; font-style: italic;">data</span> -= p.grad * learning_rate
            p.grad.zero_()   <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">reset gradient</span>
</pre>
</div>

 <p>
We'll also calculate an accuracy metric for each epoch so that we can
observe that the accuracy is improving with each successive epoch.
</p>

 <p>
First we calculate the metric for each batch:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">batch_accuracy</span>(predictions, targets):
     <span style="font-weight: bold; font-style: italic;">predictions</span> = sigmoid(predictions)
     <span style="font-weight: bold; font-style: italic;">correct</span> = (predictions > 0.5) == targets
     <span style="font-weight: bold;">return</span> correct. <span style="font-weight: bold;">float</span>().mean()
</pre>
</div>

 <p>
Then average it for the entire epoch (rounded off to 4 decimal places):
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">validate_epoch</span>(model):
     <span style="font-weight: bold; font-style: italic;">accuracies</span> = [batch_accuracy(model(x_batch), y_batch)  <span style="font-weight: bold;">for</span> x_batch, y_batch  <span style="font-weight: bold;">in</span> valid_dl]
     <span style="font-weight: bold;">return</span>  <span style="font-weight: bold;">round</span>(torch.stack(accuracies).mean().item(), 4)
</pre>
</div>

 <p>
We can now see the model performance for the entire epoch. We define a
procedure that takes in the model, parameters, learning rate and
number of epochs and prints out the accuracy:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">learn</span>(model, params, learning_rate, epochs):
     <span style="font-weight: bold;">for</span> i  <span style="font-weight: bold;">in</span>  <span style="font-weight: bold;">range</span>(epochs):
        train_epoch(model, learning_rate, params)
         <span style="font-weight: bold;">print</span>(validate_epoch(model))
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">learn(linear, params=(weights, bias), learning_rate=1.0, epochs=20)
</pre>
</div>

 <pre class="example" id="org7e14f36">
0.6142
0.7987
0.8749
0.9096
0.9252
0.9301
0.9365
0.9413
0.9438
0.9462
0.9496
0.9511
0.954
0.954
0.956
0.9565
0.9574
0.9613
0.9618
0.9628
</pre>

 <p>
We see that the accuracy gradually improves to approx 97%.
</p>
</div>
 <div id="outline-container-orgff4f698" class="outline-3">
 <h3 id="orgff4f698"> <span class="section-number-3">3.1.</span> Pytorch/FastAI conveniences</h3>
 <div class="outline-text-3" id="text-3-1">
 <p>
Pytorch provides some handy functionality that we can use to simplify
the process above.   <code>nn.Linear</code> will combine what  <code>init_params</code> and
 <code>linear</code> do together:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">linear_model</span> = nn.Linear(28 * 28, 1)
 <span style="font-weight: bold; font-style: italic;">weights</span>,  <span style="font-weight: bold; font-style: italic;">bias</span> = linear_model.parameters()
weights.shape, bias.shape
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">torch.Size</td>
 <td class="org-left">((1 784))</td>
 <td class="org-left">torch.Size</td>
 <td class="org-left">((1))</td>
</tr></tbody></table> <p>
We can also create an optimizer class that will optimize parameters
using an interface that resembles pytorch's:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">class</span>  <span style="font-weight: bold; text-decoration: underline;">BasicOptimizer</span>:
     <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">__init__</span>( <span style="font-weight: bold;">self</span>, params, learning_rate):
         <span style="font-weight: bold;">self</span>. <span style="font-weight: bold; font-style: italic;">params</span> =  <span style="font-weight: bold;">list</span>(params)
         <span style="font-weight: bold;">self</span>. <span style="font-weight: bold; font-style: italic;">learning_rate</span> = learning_rate

     <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">step</span>( <span style="font-weight: bold;">self</span>, *args, **kwargs):
         <span style="font-weight: bold;">for</span> p  <span style="font-weight: bold;">in</span>  <span style="font-weight: bold;">self</span>.params:
            p. <span style="font-weight: bold; font-style: italic;">data</span> -= p.grad.data *  <span style="font-weight: bold;">self</span>.learning_rate

     <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">zero_grad</span>( <span style="font-weight: bold;">self</span>, *args, **kwargs):
         <span style="font-weight: bold;">for</span> p  <span style="font-weight: bold;">in</span>  <span style="font-weight: bold;">self</span>.params:
            p. <span style="font-weight: bold; font-style: italic;">grad</span> =  <span style="font-weight: bold; text-decoration: underline;">None</span>
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">learning_rate</span> = 1.0
 <span style="font-weight: bold; font-style: italic;">optimizer</span> = BasicOptimizer(linear_model.parameters(), learning_rate)
</pre>
</div>

 <p>
We can re-write  <code>train_epoch</code> and  <code>learn</code> to use the new optimizer:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">train_epoch</span>(model):
     <span style="font-weight: bold;">for</span> x, y  <span style="font-weight: bold;">in</span> train_dl:
        calculate_gradients(x, y, model)
        optimizer.step()
        optimizer.zero_grad()

 <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">learn</span>(model, epochs):
     <span style="font-weight: bold;">for</span> i  <span style="font-weight: bold;">in</span>  <span style="font-weight: bold;">range</span>(epochs):
        train_epoch(model)
         <span style="font-weight: bold;">print</span>(validate_epoch(model))
</pre>
</div>

 <p>
We should get similar results to the previous run:
</p>

 <div class="org-src-container">
 <pre class="src src-python">learn(linear_model, epochs=20)
</pre>
</div>

 <pre class="example" id="org2984d88">
0.4932
0.8354
0.8418
0.9116
0.9331
0.9473
0.9555
0.9619
0.9658
0.9668
0.9687
0.9707
0.9731
0.9746
0.9761
0.977
0.9775
0.9775
0.978
0.9785
</pre>

 <p>
FastAI provides the class  <code>SGD</code> that does the same thing as  <code>BasicOptimizer</code>:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">linear_model</span> = nn.Linear(28 * 28, 1)
 <span style="font-weight: bold; font-style: italic;">optimizer</span> = SGD(linear_model.parameters(), learning_rate)
learn(linear_model, epochs=20)
</pre>
</div>

 <pre class="example" id="org0aef132">
0.4932
0.8462
0.8262
0.9106
0.9346
0.9463
0.9555
0.9619
0.9658
0.9673
0.9702
0.9717
0.9731
0.9751
0.9756
0.977
0.9775
0.978
0.978
0.9785
</pre>

 <p>
Fast AI also provides a  <code>Learner.fit</code> method which does the same thing
as our  <code>learn</code>. To use it, we combine the training and validation
dataloaders using a  <code>DataLoaders</code> object:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">dataloaders</span> = DataLoaders(train_dl, valid_dl)

 <span style="font-weight: bold; font-style: italic;">learner</span> = Learner(
    dataloaders,
    nn.Linear(28 * 28, 1),
    opt_func=SGD,
    loss_func=mnist_loss,
    metrics=batch_accuracy,
)
learner.fit(20, lr=learning_rate)
</pre>
</div>

 <pre class="example" id="org2102a06">
█epoch     train_loss  valid_loss  batch_accuracy  time    
█Epoch 1/20 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 1/20 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 1/20 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.4606]Epoch 1/20 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.2281]Epoch 1/20 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.1506]Epoch 1/20 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.1118]Epoch 1/20 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.6246]Epoch 1/20 :                                                                                    Epoch 1/20 :                                                                                    █Epoch 1/20 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 1/20 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 1/20 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.6365]Epoch 1/20 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.6365]Epoch 1/20 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.6365]Epoch 1/20 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.6365]Epoch 1/20 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.6365]Epoch 1/20 :                                                                                  Epoch 1/20 :                                                                                  0         0.636500    0.503388    0.495584        00:00     
█Epoch 2/20 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 2/20 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 2/20 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.6165]Epoch 2/20 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.5973]Epoch 2/20 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.5790]Epoch 2/20 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.5613]Epoch 2/20 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.4990]Epoch 2/20 :                                                                                    Epoch 2/20 :                                                                                    █Epoch 2/20 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 2/20 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 2/20 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.4876]Epoch 2/20 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.4876]Epoch 2/20 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.4876]Epoch 2/20 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.4876]Epoch 2/20 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.4876]Epoch 2/20 :                                                                                  Epoch 2/20 :                                                                                  1         0.487571    0.209592    0.817468        00:00     

--snipped--

█Epoch 19/20 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 19/20 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 19/20 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.0150]Epoch 19/20 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.0157]Epoch 19/20 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.0161]Epoch 19/20 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.0163]Epoch 19/20 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.0147]Epoch 19/20 :                                                                                    Epoch 19/20 :                                                                                    █Epoch 19/20 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 19/20 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 19/20 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.0145]Epoch 19/20 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.0145]Epoch 19/20 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.0145]Epoch 19/20 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.0145]Epoch 19/20 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.0145]Epoch 19/20 :                                                                                  Epoch 19/20 :                                                                                  18        0.014492    0.026389    0.977920        00:00     
█Epoch 20/20 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 20/20 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 20/20 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.0148]Epoch 20/20 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.0155]Epoch 20/20 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.0158]Epoch 20/20 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.0161]Epoch 20/20 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.0146]Epoch 20/20 :                                                                                    Epoch 20/20 :                                                                                    █Epoch 20/20 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 20/20 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 20/20 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.0143]Epoch 20/20 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.0143]Epoch 20/20 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.0143]Epoch 20/20 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.0143]Epoch 20/20 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.0143]Epoch 20/20 :                                                                                  Epoch 20/20 :                                                                                  19        0.014336    0.025804    0.978410        00:00
</pre>

 <p>
We can now upgrade our model from a linear function to a simple neural
network of two linear layers separated by a non-linearity.
</p>

 <p>
The composition of one or more linear functions results in another
linear function, but we need the linear functions decoupled from each
other to be able to model more complex patterns, hence the use of a
non-linearity.
</p>

 <p>
The non-linearity in this case is the rectified linear unit which when
given an input tensor X, outputs max(X, 0), meaning any values less
than 0 are replaced by 0.
</p>

 <p>
The first layer outputs 20 activations. The second one takes the 20
inputs and produces one activation:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">simple_net</span> = nn.Sequential(
    nn.Linear(28 * 28, 20),
    nn.ReLU(),
    nn.Linear(20, 1)
)
</pre>
</div>

 <p>
Being a deeper network, we can use a lower learning rate and more epochs:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">learner</span> = Learner(
    dataloaders, simple_net, opt_func=SGD, loss_func=mnist_loss, metrics=batch_accuracy
)
learner.fit(40, 0.1)
</pre>
</div>

 <pre class="example" id="orgd7438fa">
█epoch     train_loss  valid_loss  batch_accuracy  time    
█Epoch 1/40 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 1/40 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 1/40 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.4877]Epoch 1/40 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.4668]Epoch 1/40 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.4477]Epoch 1/40 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.4263]Epoch 1/40 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.3334]Epoch 1/40 :                                                                                    Epoch 1/40 :                                                                                    █Epoch 1/40 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 1/40 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 1/40 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.3246]Epoch 1/40 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.3246]Epoch 1/40 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.3246]Epoch 1/40 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.3246]Epoch 1/40 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.3246]Epoch 1/40 :                                                                                  Epoch 1/40 :                                                                                  0         0.324595    0.405516    0.509814        00:00     
█Epoch 2/40 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 2/40 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 2/40 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.3389]Epoch 2/40 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.3478]Epoch 2/40 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.3479]Epoch 2/40 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.3426]Epoch 2/40 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.1526]Epoch 2/40 :                                                                                    Epoch 2/40 :                                                                                    █Epoch 2/40 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 2/40 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 2/40 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.1494]Epoch 2/40 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.1494]Epoch 2/40 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.1494]Epoch 2/40 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.1494]Epoch 2/40 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.1494]Epoch 2/40 :                                                                                  Epoch 2/40 :                                                                                  1         0.149380    0.234102    0.797841        00:00     

--snipped--

█Epoch 39/40 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 39/40 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 39/40 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.0147]Epoch 39/40 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.0152]Epoch 39/40 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.0153]Epoch 39/40 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.0154]Epoch 39/40 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.0148]Epoch 39/40 :                                                                                    Epoch 39/40 :                                                                                    █Epoch 39/40 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 39/40 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 39/40 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.0146]Epoch 39/40 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.0146]Epoch 39/40 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.0146]Epoch 39/40 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.0146]Epoch 39/40 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.0146]Epoch 39/40 :                                                                                  Epoch 39/40 :                                                                                  38        0.014574    0.020761    0.982336        00:00     
█Epoch 40/40 : |-------------------------------------------| 0.00% [0/49 00:00<?]Epoch 40/40 : |-------------------------------------------| 2.04% [1/49 00:00<00:00]Epoch 40/40 : |█------------------------------------------| 4.08% [2/49 00:00<00:00... 0.0146]Epoch 40/40 : |██-----------------------------------------| 6.12% [3/49 00:00<00:00... 0.0150]Epoch 40/40 : |███----------------------------------------| 8.16% [4/49 00:00<00:00... 0.0152]Epoch 40/40 : |████---------------------------------------| 10.20% [5/49 00:00<00:00... 0.0153]Epoch 40/40 : |███████████████████████████████████████████| 100.00% [49/49 00:00<00:00... 0.0147]Epoch 40/40 :                                                                                    Epoch 40/40 :                                                                                    █Epoch 40/40 : |-------------------------------------------| 0.00% [0/8 00:00<?]Epoch 40/40 : |█████--------------------------------------| 12.50% [1/8 00:00<00:00]Epoch 40/40 : |██████████---------------------------------| 25.00% [2/8 00:00<00:00... 0.0145]Epoch 40/40 : |████████████████---------------------------| 37.50% [3/8 00:00<00:00... 0.0145]Epoch 40/40 : |█████████████████████----------------------| 50.00% [4/8 00:00<00:00... 0.0145]Epoch 40/40 : |██████████████████████████-----------------| 62.50% [5/8 00:00<00:00... 0.0145]Epoch 40/40 : |███████████████████████████████████████████| 100.00% [8/8 00:00<00:00... 0.0145]Epoch 40/40 :                                                                                  Epoch 40/40 :                                                                                  39        0.014454    0.020628    0.982336        00:00
</pre>
</div>
</div>
</div>
 <div id="outline-container-org8ca759e" class="outline-2">
 <h2 id="org8ca759e"> <span class="section-number-2">4.</span> Conclusion</h2>
 <div class="outline-text-2" id="text-4">
 <p>
At this point we have:
</p>
 <ul class="org-ul"> <li>According to the universal approximation theorem, a function that
can approximate any problem to any level of accuracy given the right
parameters</li>
 <li>A method of finding the correct parameters via stochastic gradient
descent.</li>
</ul></div>
</div>
</div>]]></content>
  <link href="https://dataai.ng/model_from_scratch_mnist.html"/>
  <id>https://dataai.ng/model_from_scratch_mnist.html</id>
  <updated>2026-07-28T14:04:00+03:00</updated>
</entry>
<entry>
  <title>From OneR to Random Forests: A Decision Trees from Scratch Approach</title>
  <author><name>KE Programmer</name></author>
  <summary>Published on Jun 15, 2023 10:00 by KE Programmer</summary>
  <content type="html"><![CDATA[<div id="content" class="content">
 <h1 class="title">From OneR to Random Forests: A Decision Trees from Scratch Approach
 <br></br> <span class="subtitle">Published on Jun 15, 2023 10:00 by KE Programmer</span>
</h1>
 <div id="table-of-contents" role="doc-toc">
 <h2>Table of Contents</h2>
 <div id="text-table-of-contents" role="doc-toc">
 <ul> <li> <a href="#orgd64bd1d">1. Introduction</a></li>
 <li> <a href="#orgec13020">2. Setup</a></li>
 <li> <a href="#org1e3a0dc">3. Data Processing</a>
 <ul> <li> <a href="#org0dcfb33">3.1. Load the Data</a></li>
 <li> <a href="#org632358b">3.2. Explore the Data</a></li>
 <li> <a href="#org3d426d7">3.3. Data Cleaning</a></li>
</ul></li>
 <li> <a href="#org515fb27">4. Binary Splits</a>
 <ul> <li> <a href="#org79d8fbf">4.1. The Score Function</a></li>
 <li> <a href="#orgb362bb9">4.2. Finding the Best Split</a></li>
</ul></li>
 <li> <a href="#org2801d7b">5. OneR (One Rule) Classifier</a>
 <ul> <li> <a href="#org4297cd0">5.1. Evaluating the OneR Model</a></li>
</ul></li>
 <li> <a href="#org5e53168">6. Decision Tree</a>
 <ul> <li> <a href="#org4373d12">6.1. Using  <code>sklearn</code>'s  <code>DecisionTreeClassifier</code></a></li>
 <li> <a href="#org13e15e1">6.2. Visualizing the Tree</a></li>
 <li> <a href="#orgf1fbedc">6.3. Evaluating the Decision Tree</a></li>
 <li> <a href="#org9456042">6.4. A Larger Decision Tree</a></li>
</ul></li>
 <li> <a href="#org3c5bf3c">7. Random Forests</a>
 <ul> <li> <a href="#org797f901">7.1. From Scratch: Bagging Multiple Trees</a></li>
 <li> <a href="#orgbd10eeb">7.2. Using  <code>sklearn</code>'s  <code>RandomForestClassifier</code></a></li>
 <li> <a href="#orgf1c9743">7.3. Feature Importance</a></li>
</ul></li>
 <li> <a href="#org6c5eecb">8. Conclusion</a></li>
</ul></div>
</div>
 <div id="outline-container-orgd64bd1d" class="outline-2">
 <h2 id="orgd64bd1d"> <span class="section-number-2">1.</span> Introduction</h2>
 <div class="outline-text-2" id="text-1">
 <p>
A decision trees from-scratch treatment for tabular data using the
 <a href="https://www.kaggle.com/competitions/icr-identify-age-related-conditions">Identify Age-Related Conditions</a> Kaggle competition.
</p>

 <p>
We'll start with a very simple model, the OneR (One Rule) classifier,
that makes predictions based on a single feature. We'll improve it
step by step by performing splits on several features, implementing a
decision tree, and finally combining many trees into a random forest.
</p>
</div>
</div>
 <div id="outline-container-orgec13020" class="outline-2">
 <h2 id="orgec13020"> <span class="section-number-2">2.</span> Setup</h2>
 <div class="outline-text-2" id="text-2">
 <p>
Start by importing commonly used modules:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> fastai.imports  <span style="font-weight: bold;">import</span> *
</pre>
</div>
</div>
</div>
 <div id="outline-container-org1e3a0dc" class="outline-2">
 <h2 id="org1e3a0dc"> <span class="section-number-2">3.</span> Data Processing</h2>
 <div class="outline-text-2" id="text-3">
 <p>
Get the dataset appropriately whether we're in Kaggle or not. If in Kaggle,
it is assumed the competition dataset has been connected to the notebook.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">import</span> os
 <span style="font-weight: bold; font-style: italic;">competition_name</span> =  <span style="font-style: italic;">"icr-identify-age-related-conditions"</span>

 <span style="font-weight: bold; font-style: italic;">is_kaggle</span> = os.environ.get( <span style="font-style: italic;">'KAGGLE_KERNEL_RUN_TYPE'</span>,  <span style="font-style: italic;">''</span>)
 <span style="font-weight: bold;">if</span> is_kaggle:
     <span style="font-weight: bold; font-style: italic;">path</span> = Path(f <span style="font-style: italic;">"/kaggle/input/</span>{competition_name} <span style="font-style: italic;">"</span>)
 <span style="font-weight: bold;">else</span>:
     <span style="font-weight: bold;">import</span> zipfile, kaggle
     <span style="font-weight: bold; font-style: italic;">path</span> = Path.home() /  <span style="font-style: italic;">'.kaggle'</span> /  <span style="font-style: italic;">'input'</span> / competition_name
    kaggle.api.competition_download_cli(competition_name, path=path.parent)
    zipfile.ZipFile(f <span style="font-style: italic;">'</span>{path} <span style="font-style: italic;">.zip'</span>).extractall(path)
</pre>
</div>
</div>
 <div id="outline-container-org0dcfb33" class="outline-3">
 <h3 id="org0dcfb33"> <span class="section-number-3">3.1.</span> Load the Data</h3>
 <div class="outline-text-3" id="text-3-1">
 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">df_train</span> = pd.read_csv(f <span style="font-style: italic;">'</span>{path} <span style="font-style: italic;">/train.csv'</span>)
 <span style="font-weight: bold; font-style: italic;">df_test</span> = pd.read_csv(f <span style="font-style: italic;">'</span>{path} <span style="font-style: italic;">/test.csv'</span>)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">df_train.head()
</pre>
</div>

 <pre class="example">
             Id        AB          AF          AH  ...         GH         GI         GL  Class
0  000ff2bfdfe9  0.209377  3109.03329   85.200147  ...  22.136229  69.834944   0.120343      1
1  007255e47698  0.145282   978.76416   85.200147  ...  29.135430  32.131996  21.978000      0
2  013f2bd269f5  0.470030  2635.10654   85.200147  ...  28.022851  35.192676   0.196941      0
3  043ac50845d5  0.252107  3819.65177  120.201618  ...  39.948656  90.493248   0.155829      0
4  044fb8a146ec  0.380297  3733.04844   85.200147  ...  45.381316  36.262628   0.096614      1

[5 rows x 58 columns]
</pre>
</div>
</div>
 <div id="outline-container-org632358b" class="outline-3">
 <h3 id="org632358b"> <span class="section-number-3">3.2.</span> Explore the Data</h3>
 <div class="outline-text-3" id="text-3-2">
 <div class="org-src-container">
 <pre class="src src-python">df_train.describe()
</pre>
</div>

 <pre class="example" id="org0166152">
               AB            AF           AH  ...          GI          GL       Class
count  617.000000    617.000000   617.000000  ...  617.000000  616.000000  617.000000
mean     0.477149   3502.013221   118.624513  ...   50.584437    8.530961    0.175041
std      0.468388   2300.322717   127.838950  ...   36.266251   10.327010    0.380310
min      0.081187    192.593280    85.200147  ...    0.897628    0.001129    0.000000
25%      0.252107   2197.345480    85.200147  ...   23.011684    0.124392    0.000000
50%      0.354659   3120.318960    85.200147  ...   41.007968    0.337827    0.000000
75%      0.559763   4361.637390   113.739540  ...   67.931664   21.978000    0.000000
max      6.161666  28688.187660  1910.123198  ...  191.194764   21.978000    1.000000

[8 rows x 56 columns]
</pre>

 <p>
We can see that the mean of the dependent column  <code>Class</code> is much
closer to zero than one. This means that observations with a positive
diagnosis are a smaller proportion of the training data. We can
confirm this by plotting a pie chart for column  <code>Class</code>.
</p>


 <div class="org-src-container">
 <pre class="src src-python">df_train.Class.value_counts().plot.pie()
</pre>
</div>


 <div id="orga539cfd" class="figure">
 <p> <img src="img/icr_class_pie.png" alt="icr_class_pie.png"></img></p>
</div>

 <p>
We also check for null values:
</p>

 <div class="org-src-container">
 <pre class="src src-python">df_train.isna(). <span style="font-weight: bold;">sum</span>()
</pre>
</div>

 <pre class="example" id="org86623c8">
Id        0
AB        0
AF        0
AH        0
AM        0
AR        0
AX        0
AY        0
AZ        0
BC        0
BD        0
BN        0
BP        0
BQ       60
BR        0
BZ        0
CB        2
CC        3
CD        0
CF        0
CH        0
CL        0
CR        0
CS        0
CU        0
CW        0
DA        0
DE        0
DF        0
DH        0
DI        0
DL        0
DN        0
DU        1
DV        0
DY        0
EB        0
EE        0
EG        0
EH        0
EJ        0
EL       60
EP        0
EU        0
FC        1
FD        0
FE        0
FI        0
FL        1
FR        0
FS        2
GB        0
GE        0
GF        0
GH        0
GI        0
GL        1
Class     0
dtype: int64
</pre>

 <div class="org-src-container">
 <pre class="src src-python">df_train.info()
</pre>
</div>

 <pre class="example" id="orgfcc867a">
<class 'pandas.DataFrame'>
RangeIndex: 617 entries, 0 to 616
Data columns (total 58 columns):
 #   Column  Non-Null Count  Dtype  
---  ------  --------------  -----  
 0   Id      617 non-null    str    
 1   AB      617 non-null    float64
 2   AF      617 non-null    float64
 3   AH      617 non-null    float64
 4   AM      617 non-null    float64
 5   AR      617 non-null    float64
 6   AX      617 non-null    float64
 7   AY      617 non-null    float64
 8   AZ      617 non-null    float64
 9   BC      617 non-null    float64
 10  BD      617 non-null    float64
 11  BN      617 non-null    float64
 12  BP      617 non-null    float64
 13  BQ      557 non-null    float64
 14  BR      617 non-null    float64
 15  BZ      617 non-null    float64
 16  CB      615 non-null    float64
 17  CC      614 non-null    float64
 18  CD      617 non-null    float64
 19  CF      617 non-null    float64
 20  CH      617 non-null    float64
 21  CL      617 non-null    float64
 22  CR      617 non-null    float64
 23  CS      617 non-null    float64
 24  CU      617 non-null    float64
 25  CW      617 non-null    float64
 26  DA      617 non-null    float64
 27  DE      617 non-null    float64
 28  DF      617 non-null    float64
 29  DH      617 non-null    float64
 30  DI      617 non-null    float64
 31  DL      617 non-null    float64
 32  DN      617 non-null    float64
 33  DU      616 non-null    float64
 34  DV      617 non-null    float64
 35  DY      617 non-null    float64
 36  EB      617 non-null    float64
 37  EE      617 non-null    float64
 38  EG      617 non-null    float64
 39  EH      617 non-null    float64
 40  EJ      617 non-null    str    
 41  EL      557 non-null    float64
 42  EP      617 non-null    float64
 43  EU      617 non-null    float64
 44  FC      616 non-null    float64
 45  FD      617 non-null    float64
 46  FE      617 non-null    float64
 47  FI      617 non-null    float64
 48  FL      616 non-null    float64
 49  FR      617 non-null    float64
 50  FS      615 non-null    float64
 51  GB      617 non-null    float64
 52  GE      617 non-null    float64
 53  GF      617 non-null    float64
 54  GH      617 non-null    float64
 55  GI      617 non-null    float64
 56  GL      616 non-null    float64
 57  Class   617 non-null    int64  
dtypes: float64(55), int64(1), str(2)
memory usage: 279.7 KB
</pre>
</div>
</div>
 <div id="outline-container-org3d426d7" class="outline-3">
 <h3 id="org3d426d7"> <span class="section-number-3">3.3.</span> Data Cleaning</h3>
 <div class="outline-text-3" id="text-3-3">
 <p>
On the competition's data tab, we're informed that all columns are
numeric with the exception of  <code>EJ</code>, which is categorical. We'll
replace null values with modes, and convert  <code>EJ</code> to a pandas
categorical column.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">modes</span> = df_train.mode().iloc[0]
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">process_data</span>(df):
    df.fillna(modes, inplace= <span style="font-weight: bold; text-decoration: underline;">True</span>)
     <span style="font-weight: bold; font-style: italic;">df</span>[ <span style="font-style: italic;">"EJ"</span>] = pd.Categorical(df.EJ)

process_data(df_train)
process_data(df_test)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">df_train.info()
</pre>
</div>

 <pre class="example" id="org7b36615">
<class 'pandas.DataFrame'>
RangeIndex: 617 entries, 0 to 616
Data columns (total 58 columns):
 #   Column  Non-Null Count  Dtype   
---  ------  --------------  -----   
 0   Id      617 non-null    str     
 1   AB      617 non-null    float64 
 2   AF      617 non-null    float64 
 3   AH      617 non-null    float64 
 4   AM      617 non-null    float64 
 5   AR      617 non-null    float64 
 6   AX      617 non-null    float64 
 7   AY      617 non-null    float64 
 8   AZ      617 non-null    float64 
 9   BC      617 non-null    float64 
 10  BD      617 non-null    float64 
 11  BN      617 non-null    float64 
 12  BP      617 non-null    float64 
 13  BQ      617 non-null    float64 
 14  BR      617 non-null    float64 
 15  BZ      617 non-null    float64 
 16  CB      617 non-null    float64 
 17  CC      617 non-null    float64 
 18  CD      617 non-null    float64 
 19  CF      617 non-null    float64 
 20  CH      617 non-null    float64 
 21  CL      617 non-null    float64 
 22  CR      617 non-null    float64 
 23  CS      617 non-null    float64 
 24  CU      617 non-null    float64 
 25  CW      617 non-null    float64 
 26  DA      617 non-null    float64 
 27  DE      617 non-null    float64 
 28  DF      617 non-null    float64 
 29  DH      617 non-null    float64 
 30  DI      617 non-null    float64 
 31  DL      617 non-null    float64 
 32  DN      617 non-null    float64 
 33  DU      617 non-null    float64 
 34  DV      617 non-null    float64 
 35  DY      617 non-null    float64 
 36  EB      617 non-null    float64 
 37  EE      617 non-null    float64 
 38  EG      617 non-null    float64 
 39  EH      617 non-null    float64 
 40  EJ      617 non-null    category
 41  EL      617 non-null    float64 
 42  EP      617 non-null    float64 
 43  EU      617 non-null    float64 
 44  FC      617 non-null    float64 
 45  FD      617 non-null    float64 
 46  FE      617 non-null    float64 
 47  FI      617 non-null    float64 
 48  FL      617 non-null    float64 
 49  FR      617 non-null    float64 
 50  FS      617 non-null    float64 
 51  GB      617 non-null    float64 
 52  GE      617 non-null    float64 
 53  GF      617 non-null    float64 
 54  GH      617 non-null    float64 
 55  GI      617 non-null    float64 
 56  GL      617 non-null    float64 
 57  Class   617 non-null    int64   
dtypes: category(1), float64(55), int64(1), str(1)
memory usage: 275.5 KB
</pre>

 <p>
EJ now has a category column type.
</p>

 <p>
A decision tree only requires that the column values can be ordered
numerically. For the categorical column  <code>EJ</code>, we'll use the underlying
categorical codes as its values.
</p>

 <div class="org-src-container">
 <pre class="src src-python">df_train.EJ.head()
</pre>
</div>

 <pre class="example">
0    B
1    A
2    B
3    B
4    B
Name: EJ, dtype: category
Categories (2, str): ['A', 'B']
</pre>


 <div class="org-src-container">
 <pre class="src src-python">df_train.EJ.cat.codes.head()
</pre>
</div>

 <pre class="example">
0    1
1    0
2    1
3    1
4    1
dtype: int8
</pre>



 <p>
Segregate the categorical, numeric and dependent variables:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">categoricals</span> = [ <span style="font-style: italic;">"EJ"</span>]
 <span style="font-weight: bold; font-style: italic;">dependent</span> =  <span style="font-style: italic;">"Class"</span>
 <span style="font-weight: bold; font-style: italic;">conts</span> = [column  <span style="font-weight: bold;">for</span> column  <span style="font-weight: bold;">in</span> df_train.columns
          <span style="font-weight: bold;">if</span>  <span style="font-weight: bold;">not</span> column  <span style="font-weight: bold;">in</span> categoricals + [dependent] + [ <span style="font-style: italic;">"Id"</span>]]
</pre>
</div>
</div>
</div>
</div>
 <div id="outline-container-org515fb27" class="outline-2">
 <h2 id="org515fb27"> <span class="section-number-2">4.</span> Binary Splits</h2>
 <div class="outline-text-2" id="text-4">
 <p>
A decision tree is built on binary splits, i.e. using the value of a column
to split the rows into two groups. Let's use a barplot and countplot to
analyse how splitting on  <code>EJ</code> relates to the diagnosed class.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">import</span> seaborn  <span style="font-weight: bold;">as</span> sns

 <span style="font-weight: bold; font-style: italic;">fig</span>,  <span style="font-weight: bold; font-style: italic;">axs</span> = plt.subplots(1, 2, figsize=(11, 5))
sns.barplot(data=df_train, y=dependent, x= <span style="font-style: italic;">"EJ"</span>, ax=axs[0])\
    . <span style="font-weight: bold;">set</span>(title= <span style="font-style: italic;">"Positive Diagnosis Rate"</span>)
sns.countplot(data=df_train, x= <span style="font-style: italic;">"EJ"</span>, ax=axs[1])\
    . <span style="font-weight: bold;">set</span>(title= <span style="font-style: italic;">"Histogram"</span>)
</pre>
</div>


 <div id="orgb62a4f4" class="figure">
 <p> <img src="img/icr_ej_split.png" alt="icr_ej_split.png"></img></p>
</div>

 <p>
We have a higher positivity rate for category B (~20%) than category A
(~13%). We also have a much higher proportion of observations with category
B (~400) than category A (~230).
</p>

 <p>
We can also do a split based on a continuous column. We'll go with column
 <code>AB</code> just for demonstration. We use a boxplot to compare the averages of
both positive and negative diagnosis based on the trait  <code>AB</code> and a density
plot to visualize the distribution of observations on  <code>AB</code>.
</p>

 <div class="org-src-container">
 <pre class="src src-python">fig, ( <span style="font-weight: bold; font-style: italic;">ax1</span>,  <span style="font-weight: bold; font-style: italic;">ax2</span>) = plt.subplots(nrows=1, ncols=2, figsize=(11, 5))
sns.boxenplot(data=df_train, x=dependent, y= <span style="font-style: italic;">"AB"</span>, ax=ax1)
sns.kdeplot(data=df_train, x= <span style="font-style: italic;">"AB"</span>, ax=ax2)
</pre>
</div>


 <div id="org3204004" class="figure">
 <p> <img src="img/icr_ab_distribution.png" alt="icr_ab_distribution.png"></img></p>
</div>
</div>
 <div id="outline-container-org79d8fbf" class="outline-3">
 <h3 id="org79d8fbf"> <span class="section-number-3">4.1.</span> The Score Function</h3>
 <div class="outline-text-3" id="text-4-1">
 <p>
Since we have a large number of columns, it would be tedious to plot
each of them and figure out how well it partitions between positive
and negative diagnosis. We can create a function that helps us quickly
evaluate different splits by calculating a measure of impurity.
</p>

 <p>
The key idea in determining a good split is that the it the dependent
variable is as homogenous as possible within each group.
</p>

 <p>
The lower the standard deviation of the dependent variable in a
group, the more homogenous it is.
</p>

 <p>
Next we multiple the std deviation with the group size, so that a
larger group contributes more to the score, before normalizing the
score using the total number of observations.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">score</span>(col, y, split_value):
     <span style="font-weight: bold; font-style: italic;">lhs</span> = col <= split_value
     <span style="font-weight: bold;">return</span> (_side_score(lhs, y) + _side_score(~lhs, y))/ <span style="font-weight: bold;">len</span>(y)

 <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">_side_score</span>(side, y):
     <span style="font-weight: bold; font-style: italic;">count</span> = side. <span style="font-weight: bold;">sum</span>()
     <span style="font-weight: bold;">if</span> count <=1:  <span style="font-weight: bold;">return</span> 0
     <span style="font-weight: bold;">return</span> y[side].std() * count
</pre>
</div>

 <p>
For instance, the score based on value 0.5 for column  <code>AB</code>:
</p>

 <div class="org-src-container">
 <pre class="src src-python">score(df_train.AB, df_train[dependent], 0.5)
</pre>
</div>

 <pre class="example">
0.36017711175753037
</pre>


 <p>
For the categorical column, we'd need to replace the column string
values with their underlying codes:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">df_train_n</span> = df_train.copy()
 <span style="font-weight: bold; font-style: italic;">df_train_n</span>[categoricals] = df_train_n[categoricals].apply( <span style="font-weight: bold;">lambda</span> x: x.cat.codes)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">df_train_n[categoricals].head()
</pre>
</div>

 <pre class="example">
   EJ
0   1
1   0
2   1
3   1
4   1
</pre>


 <p>
We then calculate the score using <= 1 as our split:
</p>

 <div class="org-src-container">
 <pre class="src src-python">score(df_train_n.EJ, df_train_n[dependent], 1)
</pre>
</div>

 <pre class="example">
0.3803100751041243
</pre>
</div>
</div>
 <div id="outline-container-orgb362bb9" class="outline-3">
 <h3 id="orgb362bb9"> <span class="section-number-3">4.2.</span> Finding the Best Split</h3>
 <div class="outline-text-3" id="text-4-2">
 <p>
To help us find the best split point, we'll iterate through the
columns, and for each, we iterate through its unique values to find
the best split point for the column.
</p>

 <p>
For example, to find the best split point for column  <code>AB</code>:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">col</span> = df_train_n[ <span style="font-style: italic;">"AB"</span>]
 <span style="font-weight: bold; font-style: italic;">y</span> = df_train_n[dependent]
 <span style="font-weight: bold; font-style: italic;">uniques</span> = col.unique()
uniques.sort()
uniques[:50]
</pre>
</div>

 <pre class="example">
array([0.081187 , 0.08546  , 0.098279 , 0.102552 , 0.111098 , 0.119644 ,
       0.1217805, 0.1303265, 0.132463 , 0.136736 , 0.141009 , 0.145282 ,
       0.149555 , 0.153828 , 0.1559645, 0.158101 , 0.1602375, 0.162374 ,
       0.166647 , 0.17092  , 0.175193 , 0.179466 , 0.183739 , 0.1858755,
       0.188012 , 0.1901485, 0.192285 , 0.196558 , 0.200831 , 0.2029675,
       0.205104 , 0.209377 , 0.21365  , 0.217923 , 0.222196 , 0.2243325,
       0.226469 , 0.230742 , 0.235015 , 0.239288 , 0.243561 , 0.247834 ,
       0.252107 , 0.25638  , 0.260653 , 0.2627895, 0.264926 , 0.269199 ,
       0.2713355, 0.273472 ])
</pre>



 <p>
Get the best split for  <code>AB</code>:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">scores</span> = np.array([score(col, y, split)  <span style="font-weight: bold;">for</span> split  <span style="font-weight: bold;">in</span> uniques])
uniques[scores.argmin()]
</pre>
</div>

 <pre class="example">
0.410208
</pre>


 <p>
Create a function that gets the best split for for any column:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">min_column</span>(df, col_name):
     <span style="font-weight: bold; font-style: italic;">col</span>,  <span style="font-weight: bold; font-style: italic;">y</span> = df[col_name], df[dependent]
     <span style="font-weight: bold; font-style: italic;">uniques</span> = col.unique()
     <span style="font-weight: bold; font-style: italic;">scores</span> = np.array([score(col, y, split)  <span style="font-weight: bold;">for</span> split  <span style="font-weight: bold;">in</span> uniques])
     <span style="font-weight: bold; font-style: italic;">index</span> = scores.argmin()
     <span style="font-weight: bold;">return</span> uniques[index], scores[index]
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">min_column(df_train_n,  <span style="font-style: italic;">"AB"</span>)
</pre>
</div>

 <p>
We then try the same for the categorical variable:
</p>

 <div class="org-src-container">
 <pre class="src src-python">min_column(df_train_n,  <span style="font-style: italic;">"EJ"</span>)
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">np.int8</td>
 <td class="org-left">(0)</td>
 <td class="org-left">np.float64</td>
 <td class="org-left">(0.3773339803468088)</td>
</tr></tbody></table> <p>
We then calculate the best split points for each of the columns to
find the best split overall:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">columns</span> = conts + categoricals
 <span style="font-weight: bold; font-style: italic;">splits</span> = {col: min_column(df_train_n, col)  <span style="font-weight: bold;">for</span> col  <span style="font-weight: bold;">in</span> columns}
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">splits
</pre>
</div>

 <pre class="example">
{'AB': (np.float64(0.410208), np.float64(0.3561906829723693)), 'AF': (np.float64(2808.64232), np.float64(0.35583548065304915)), 'AH': (np.float64(193.801377), np.float64(0.3761985927636321)), 'AM': (np.float64(149.318758), np.float64(0.367338331639673)), 'AR': (np.float64(16.327194), np.float64(0.3678281215243884)), 'AX': (np.float64(17.877462), np.float64(0.37677322461475915)), 'AY': (np.float64(0.6019965), np.float64(0.37677322461475915)), 'AZ': (np.float64(10.971782), np.float64(0.3782810033906922)), 'BC': (np.float64(13.500788), np.float64(0.35968813799570454)), 'BD ': (np.float64(12083.34891), np.float64(0.3749922782916859)), 'BN': (np.float64(21.186), np.float64(0.3651185274145938)), 'BP': (np.float64(196.710795), np.float64(0.3725348222558456)), 'BQ': (np.float64(115.695865), np.float64(0.3668876611761174)), 'BR': (np.float64(3466.745415), np.float64(0.374094640781585)), 'BZ': (np.float64(2885.319798), np.float64(0.376211243844383)), 'CB': (np.float64(13.32695), np.float64(0.37775055805021784)), 'CC': (np.float64(0.5478777), np.float64(0.3665150807004323)), 'CD ': (np.float64(85.955376), np.float64(0.3635292346687648)), 'CF': (np.float64(1.8504485), np.float64(0.367406050093466)), 'CH': (np.float64(0.016318), np.float64(0.3782672405025747)), 'CL': (np.float64(1.24754), np.float64(0.37657306471010377)), 'CR': (np.float64(0.527325), np.float64(0.35228611051878406)), 'CS': (np.float64(62.2516675), np.float64(0.37277362672289577)), 'CU': (np.float64(1.274427), np.float64(0.37534114864714596)), 'CW ': (np.float64(35.67944), np.float64(0.3772898765041537)), 'DA': (np.float64(27.36564), np.float64(0.353137425790351)), 'DE': (np.float64(149.18453), np.float64(0.3625701653246753)), 'DF': (np.float64(0.500175), np.float64(0.36551421649454285)), 'DH': (np.float64(0.240504), np.float64(0.3671902453269601)), 'DI': (np.float64(253.8155925), np.float64(0.3550252288774408)), 'DL': (np.float64(134.1642), np.float64(0.37162996473999704)), 'DN': (np.float64(59.068544), np.float64(0.3749922782916859)), 'DU': (np.float64(2.27601), np.float64(0.3207517917154871)), 'DV': (np.float64(2.19891), np.float64(0.37819979248926816)), 'DY': (np.float64(4.474032), np.float64(0.3708655621394996)), 'EB': (np.float64(6.269316), np.float64(0.3625606667600181)), 'EE': (np.float64(1.463253), np.float64(0.36191183975348784)), 'EG': (np.float64(6845.912275), np.float64(0.37731357741606586)), 'EH': (np.float64(0.389376), np.float64(0.3593626675299252)), 'EL': (np.float64(46.05744), np.float64(0.37829464648148575)), 'EP': (np.float64(224.078075), np.float64(0.37429306818140773)), 'EU': (np.float64(8.497392), np.float64(0.3699652228963962)), 'FC': (np.float64(13.33752), np.float64(0.37580636037706155)), 'FD ': (np.float64(8.151501), np.float64(0.3607755013721559)), 'FE': (np.float64(15667.04141), np.float64(0.3658514670660312)), 'FI': (np.float64(8.9724075), np.float64(0.3657210528731738)), 'FL': (np.float64(7.925430474), np.float64(0.3436877494985317)), 'FR': (np.float64(2.73702), np.float64(0.365585424163423)), 'FS': (np.float64(0.839852), np.float64(0.37680072033409695)), 'GB': (np.float64(40.435794), np.float64(0.37726637931071944)), 'GE': (np.float64(363.134821), np.float64(0.37718575251296815)), 'GF': (np.float64(14737.27446), np.float64(0.37078665888029855)), 'GH': (np.float64(61.028121), np.float64(0.3762112438443829)), 'GI': (np.float64(117.6047), np.float64(0.3767203577225661)), 'GL': (np.float64(0.121055405), np.float64(0.3450719268397208)), 'EJ': (np.int8(0), np.float64(0.3773339803468088))}
</pre>


 <p>
It is tedious to visually identify the best split, so we automate it:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">min_score</span> = 1
 <span style="font-weight: bold;">for</span> k, v  <span style="font-weight: bold;">in</span> splits.items():
     <span style="font-weight: bold;">if</span> v[1] < min_score:
         <span style="font-weight: bold; font-style: italic;">min_score</span> = v[1]
         <span style="font-weight: bold; font-style: italic;">min_score_column</span> = k
min_score_column, splits[min_score_column]
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">DU</td>
 <td class="org-left">(np.float64 (2.27601) np.float64 (0.3207517917154871))</td>
</tr></tbody></table> <p>
According to this, the column  <code>DU</code> gives the best score at split point
 <code>2.27601</code> overall. This gives us a simple model based on a single
rule, a variant of what is called the  <a href="https://link.springer.com/article/10.1023/A:1022631118932">OneR</a> (One Rule) classifier.
</p>
</div>
</div>
</div>
 <div id="outline-container-org2801d7b" class="outline-2">
 <h2 id="org2801d7b"> <span class="section-number-2">5.</span> OneR (One Rule) Classifier</h2>
 <div class="outline-text-2" id="text-5">
 <p>
Though it's already a small dataset, we'll go through the typical
process of carving out a small validation set to evaluate the model.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> numpy  <span style="font-weight: bold;">import</span> random
 <span style="font-weight: bold;">from</span> sklearn.model_selection  <span style="font-weight: bold;">import</span> train_test_split

random.seed(42)
 <span style="font-weight: bold; font-style: italic;">model_train</span>,  <span style="font-weight: bold; font-style: italic;">model_val</span> = train_test_split(df_train_n, test_size=0.25)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">model_train.shape, model_val.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">462</td>
 <td class="org-right">58</td>
</tr> <tr> <td class="org-right">155</td>
 <td class="org-right">58</td>
</tr></tbody></table> <p>
We split each of training and validation sets into independent and
dependent variables:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">x_y</span>(df):
     <span style="font-weight: bold;">return</span> df[conts + categoricals], df[dependent]

 <span style="font-weight: bold; font-style: italic;">model_train_x</span>,  <span style="font-weight: bold; font-style: italic;">model_train_y</span> = x_y(model_train)
 <span style="font-weight: bold; font-style: italic;">model_val_x</span>,  <span style="font-weight: bold; font-style: italic;">model_val_y</span> = x_y(model_val)
</pre>
</div>

 <p>
We get the best split using our model's training set:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">one_r_splits</span> = {col: min_column(model_train, col)
                 <span style="font-weight: bold;">for</span> col  <span style="font-weight: bold;">in</span> model_train_x.columns}
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">get_best_split</span>(splits):
     <span style="font-weight: bold; font-style: italic;">min_score</span> = 1
     <span style="font-weight: bold;">for</span> k, v  <span style="font-weight: bold;">in</span> splits.items():
         <span style="font-weight: bold;">if</span> v[1] < min_score:
             <span style="font-weight: bold; font-style: italic;">min_score</span> = v[1]
             <span style="font-weight: bold; font-style: italic;">min_score_column</span> = k
     <span style="font-weight: bold;">return</span> min_score_column, splits[min_score_column]

 <span style="font-weight: bold; font-style: italic;">column</span>,  <span style="font-weight: bold; font-style: italic;">split</span> = get_best_split(one_r_splits)
column, split
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">DU</td>
 <td class="org-left">(np.float64 (2.262216) np.float64 (0.3196074664465314))</td>
</tr></tbody></table> <p>
According to this,  <code>DU</code> still remains the best split at split point
 <code>2.262216</code>.
</p>

 <p>
We visualize how splitting at this value correlates with the
dependent variable  <code>Class</code>
</p>

 <div class="org-src-container">
 <pre class="src src-python">
 <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">Create the categorical column
</span> <span style="font-weight: bold; font-style: italic;">model_train</span>[ <span style="font-style: italic;">'DU_split'</span>] = model_train[ <span style="font-style: italic;">'DU'</span>] > 2.262216

 <span style="font-weight: bold; font-style: italic;"># </span> <span style="font-weight: bold; font-style: italic;">Plot – barplot shows mean Class for each group
</span> <span style="font-weight: bold; font-style: italic;">fig</span>,  <span style="font-weight: bold; font-style: italic;">ax</span> = plt.subplots(1, 1, figsize=(11, 5))
sns.barplot(data=model_train, x= <span style="font-style: italic;">'DU_split'</span>, y= <span style="font-style: italic;">'Class'</span>, ax=ax)
plt.xlabel( <span style="font-style: italic;">'DU > 2.262216 (split)'</span>)
plt.ylabel( <span style="font-style: italic;">'Positive diagnosis rate'</span>)
</pre>
</div>

 <pre class="example">
Text(0, 0.5, 'Positive diagnosis rate')
</pre>



 <div id="orge221d3c" class="figure">
 <p> <img src="img/du_split_barplot.png" alt="du_split_barplot.png"></img></p>
</div>

 <p>
This shows that the split at 2.262216 is highly effective at
separating the two classes. The right group (> 2.262216) has a 60%
positivity rate compared to the left group at approximately 10%.
</p>

 <p>
We then use this as a simple OneR model and make predictions for the
validation set.
</p>


 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">preds</span> = model_val_x[ <span style="font-style: italic;">"DU"</span>] > 2.262216
</pre>
</div>
</div>
 <div id="outline-container-org4297cd0" class="outline-3">
 <h3 id="org4297cd0"> <span class="section-number-3">5.1.</span> Evaluating the OneR Model</h3>
 <div class="outline-text-3" id="text-5-1">
 <p>
From this we can calculate the mean absolute error to see how off the
predictions are from the actual values in the validation set:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.metrics  <span style="font-weight: bold;">import</span> mean_absolute_error
mean_absolute_error(model_val_y, preds)
</pre>
</div>

 <pre class="example">
0.13548387096774195
</pre>


 <p>
The error is relatively low, suggesting that this baseline model
decent enough for a baseline.
</p>

 <p>
The competition defines the  <a href="https://www.kaggle.com/competitions/icr-identify-age-related-conditions/overview/evaluation">metric</a> used to evaluate submissions. I've
copied the implementation from a  <a href="https://www.kaggle.com/competitions/icr-identify-age-related-conditions/discussion/410864#2265438">comment</a> in the discussion. The
implementation details of this do not matter for our purposes. We just
want to gauge the relative score as we go along.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">balanced_logarithmic_loss</span>(y_true, y_pred):
     <span style="font-weight: bold; font-style: italic;">N_1</span> = np. <span style="font-weight: bold;">sum</span>(y_true == 1, axis=0)
     <span style="font-weight: bold; font-style: italic;">N_0</span> = np. <span style="font-weight: bold;">sum</span>(y_true == 0, axis=0)
     <span style="font-weight: bold; font-style: italic;">y_pred</span> = np.maximum(np.minimum(y_pred, 1 - 1e-15), 1e-15)
     <span style="font-weight: bold; font-style: italic;">loss_numerator</span> = (- (1/N_0) * np. <span style="font-weight: bold;">sum</span>((1 - y_true) * np.log(1-y_pred))
                      - (1/N_1) * np. <span style="font-weight: bold;">sum</span>(y_true * np.log(y_pred)))
     <span style="font-weight: bold;">return</span> loss_numerator / 2

balanced_logarithmic_loss(model_val_y.to_numpy(), preds.to_numpy())
</pre>
</div>

 <pre class="example">
8.588667983985555
</pre>


 <p>
Our local score is 8.59. When submitted to the competition (after the
deadline), it had a leaderboard score of 8.56 (public) and 10.97
(private)
</p>
</div>
</div>
</div>
 <div id="outline-container-org5e53168" class="outline-2">
 <h2 id="org5e53168"> <span class="section-number-2">6.</span> Decision Tree</h2>
 <div class="outline-text-2" id="text-6">
 <p>
After having identified the best split for the training set, we can split
the data into two groups, then for each group find the next best split.
</p>

 <p>
We split our training data into two groups based on the value of  <code>DU</code> we
identified above, then find the best successive splits.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">lhs</span> = model_train[ <span style="font-style: italic;">"DU"</span>] <= 2.262216
 <span style="font-weight: bold; font-style: italic;">left_group</span> = model_train[lhs]
 <span style="font-weight: bold; font-style: italic;">right_group</span> = model_train[~lhs]
left_group.shape, right_group.shape
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-right"></col></colgroup> <tbody> <tr> <td class="org-right">396</td>
 <td class="org-right">59</td>
</tr> <tr> <td class="org-right">66</td>
 <td class="org-right">59</td>
</tr></tbody></table> <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">second_level_columns</span> = [c  <span style="font-weight: bold;">for</span> c  <span style="font-weight: bold;">in</span> (conts + categoricals)  <span style="font-weight: bold;">if</span> c !=  <span style="font-style: italic;">"DU"</span>]
 <span style="font-weight: bold; font-style: italic;">left_splits</span> = {col: min_column(left_group, col)
                <span style="font-weight: bold;">for</span> col  <span style="font-weight: bold;">in</span> second_level_columns}
 <span style="font-weight: bold; font-style: italic;">right_splits</span> = {col: min_column(right_group, col)
                 <span style="font-weight: bold;">for</span> col  <span style="font-weight: bold;">in</span> second_level_columns}
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">best_left_split</span> = get_best_split(left_splits)
 <span style="font-weight: bold; font-style: italic;">best_right_split</span> = get_best_split(right_splits)
best_left_split, best_right_split
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">AB</td>
 <td class="org-left">(np.float64 (0.363205) np.float64 (0.2205227532918278))</td>
</tr> <tr> <td class="org-left">GL</td>
 <td class="org-left">(np.float64 (0.047863636) np.float64 (0.40470758822889663))</td>
</tr></tbody></table> <p>
For our left group, column  <code>AB</code> results in the best split and  <code>GL</code> for
the right group. Combining the three splitting rules first splitting
by  <code>DU</code>, then the left group by  <code>AB</code> and the right by  <code>GL</code> results in
a decision tree.
</p>
</div>
 <div id="outline-container-org4373d12" class="outline-3">
 <h3 id="org4373d12"> <span class="section-number-3">6.1.</span> Using  <code>sklearn</code>'s  <code>DecisionTreeClassifier</code></h3>
 <div class="outline-text-3" id="text-6-1">
 <p>
Rather than rolling it out by hand, we can use sklearn's built-in Decision
Tree classifier:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.tree  <span style="font-weight: bold;">import</span> DecisionTreeClassifier, export_graphviz

 <span style="font-weight: bold; font-style: italic;">model</span> = DecisionTreeClassifier(max_leaf_nodes=4).fit(model_train_x, model_train_y)
</pre>
</div>
</div>
</div>
 <div id="outline-container-org13e15e1" class="outline-3">
 <h3 id="org13e15e1"> <span class="section-number-3">6.2.</span> Visualizing the Tree</h3>
 <div class="outline-text-3" id="text-6-2">
 <p>
We write a procedure to visualize the created tree:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">import</span> graphviz
 <span style="font-weight: bold;">import</span> re

 <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">draw_tree</span>(tree, df, size=10, ratio=0.6, precision=2, **kwargs):
     <span style="font-weight: bold; font-style: italic;">dot_format</span> = export_graphviz(
        tree, out_file= <span style="font-weight: bold; text-decoration: underline;">None</span>, feature_names=df.columns,
        filled= <span style="font-weight: bold; text-decoration: underline;">True</span>, rounded= <span style="font-weight: bold; text-decoration: underline;">True</span>, special_characters= <span style="font-weight: bold; text-decoration: underline;">True</span>,
        rotate= <span style="font-weight: bold; text-decoration: underline;">False</span>, precision=precision, **kwargs)
     <span style="font-weight: bold;">return</span> graphviz.Source(
        re.sub( <span style="font-style: italic;">'Tree {'</span>, f <span style="font-style: italic;">'Tree {{ size=</span>{size} <span style="font-style: italic;">; ratio=</span>{ratio} <span style="font-style: italic;">'</span>,
               dot_format))
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">draw_tree(model, model_train_x, size=10)
</pre>
</div>


 <div id="org322bf11" class="figure">
 <p> <img src="img/icr_decision_tree_small.svg" alt="icr_decision_tree_small.svg" class="org-svg"></img></p>
</div>

 <p>
The decision tree uses a measure of impurity called the  <i>gini index</i>. This
measures the probability that if you pick two observations from a group,
they will not have the same value for the dependent column. In the case of
perfect classification, where all observations in the group have the same
value for  <code>Class</code>, the gini index is zero.
</p>

 <p>
It is determined by subtracting the sum of squared probabilities of
each class of the prediction from 1:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">gini</span>(df, condition):
     <span style="font-weight: bold; font-style: italic;">actual</span> = df.loc[condition, dependent]
     <span style="font-weight: bold;">return</span> 1 - actual.mean()**2 - (1-actual).mean()**2
</pre>
</div>

 <p>
Simulating the split at the root node above:
</p>

 <div class="org-src-container">
 <pre class="src src-python">gini(model_train, model_train[ <span style="font-style: italic;">'DU'</span>] <= 2.28), gini(model_train, model_train[ <span style="font-style: italic;">'DU'</span>] > 2.28)
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">np.float64</td>
 <td class="org-left">(0.16940873380267307)</td>
 <td class="org-left">np.float64</td>
 <td class="org-left">(0.47061524334251614)</td>
</tr></tbody></table> <p>
Like with our OneR approach, the decision tree starts with a split on
DU as the best split, thought it uses a different value for the
split.
</p>

 <p>
The non-leaf nodes show which column was used for the split, the
split value, the gini score, the number of observations in that group
prior to the split as  <code>samples</code> and  <code>value</code> hints at the purity of that
group, showing how it is partitioned by the dependent variable.
</p>
</div>
</div>
 <div id="outline-container-orgf1fbedc" class="outline-3">
 <h3 id="orgf1fbedc"> <span class="section-number-3">6.3.</span> Evaluating the Decision Tree</h3>
 <div class="outline-text-3" id="text-6-3">
 <p>
We can calculate the mean absolute error of this decision tree:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">preds</span> = model.predict(model_val_x)
mean_absolute_error(model_val_y, preds)
</pre>
</div>

 <pre class="example">
0.11612903225806452
</pre>


 <p>
And also calculate the logloss metric:
</p>

 <div class="org-src-container">
 <pre class="src src-python">balanced_logarithmic_loss(model_val_y.to_numpy(), preds)
</pre>
</div>

 <pre class="example">
5.549265256402579
</pre>


 <p>
We see that just two additional splits have increased the accuracy on the
training data considerably.
</p>
</div>
</div>
 <div id="outline-container-org9456042" class="outline-3">
 <h3 id="org9456042"> <span class="section-number-3">6.4.</span> A Larger Decision Tree</h3>
 <div class="outline-text-3" id="text-6-4">
 <p>
We can create a bigger tree to see whether it will minimise the error
further:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">model</span> = DecisionTreeClassifier(min_samples_leaf=50)
model.fit(model_train_x, model_train_y)
draw_tree(model, model_train_x, size=12)
</pre>
</div>


 <div id="org16887f6" class="figure">
 <p> <img src="img/icr_decision_tree_large.svg" alt="icr_decision_tree_large.svg" class="org-svg"></img></p>
</div>

 <p>
By allowing the tree to have a greater number of leaf nodes, we've
allowed it to reach a grouping with a gini score of zero, which is
perfect classification. However, the larger a decision tree is, the
more it tends to overfit the training data and may not generalize well
to the test data.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">preds</span> = model.predict(model_val_x)
mean_absolute_error(model_val_y, preds),balanced_logarithmic_loss(model_val_y.to_numpy(), preds)
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-right">0.12903225806451613</td>
 <td class="org-left">np.float64</td>
 <td class="org-left">(8.450509680016193)</td>
</tr></tbody></table> <p>
The slightly higher error than the smaller tree could hint at
overfitting. When submitted to the competition, it had a score of 1.72
on the public leaderboard, an improvement over 25.97 of the OneR model.
</p>
</div>
</div>
</div>
 <div id="outline-container-org3c5bf3c" class="outline-2">
 <h2 id="org3c5bf3c"> <span class="section-number-2">7.</span> Random Forests</h2>
 <div class="outline-text-2" id="text-7">
 <p>
As mentioned previously, making the decision tree bigger has it match
the training data more closely, resulting in overfitting.
</p>

 <p>
Instead of using bigger trees, we can use  <i>more</i> trees. That's the
insight from  <a href="https://en.wikipedia.org/wiki/Leo_Breiman">Leo Breiman</a> who helped formulate the technique. By
training more trees, each trained on a random uncorrelated subset of
the training data and averaging their results, we get a better
result. This is because the average of uncorrelated errors is close to
zero.
</p>
</div>
 <div id="outline-container-org797f901" class="outline-3">
 <h3 id="org797f901"> <span class="section-number-3">7.1.</span> From Scratch: Bagging Multiple Trees</h3>
 <div class="outline-text-3" id="text-7-1">
 <p>
To demonstrate the technique, we can train a tree on a random subset of the
data:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">def</span>  <span style="font-weight: bold;">get_tree</span>(proportion=0.75):
     <span style="font-weight: bold; font-style: italic;">n</span> =  <span style="font-weight: bold;">len</span>(model_train_y)
     <span style="font-weight: bold; font-style: italic;">indexes</span> = random.choice(n,  <span style="font-weight: bold;">int</span>(n * proportion))
     <span style="font-weight: bold;">return</span> DecisionTreeClassifier(min_samples_leaf=5).fit(
        model_train_x.iloc[indexes], model_train_y.iloc[indexes]
    )
</pre>
</div>

 <p>
Now we can train as many trees as needed and average their results:
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">trees</span> = [get_tree()  <span style="font-weight: bold;">for</span> _  <span style="font-weight: bold;">in</span>  <span style="font-weight: bold;">range</span>(100)]
trees[:3]
</pre>
</div>

 <table> <colgroup> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-left">DecisionTreeClassifier</td>
 <td class="org-left">(min <sub>samples</sub> <sub>leaf</sub>=5)</td>
 <td class="org-left">DecisionTreeClassifier</td>
 <td class="org-left">(min <sub>samples</sub> <sub>leaf</sub>=5)</td>
 <td class="org-left">DecisionTreeClassifier</td>
 <td class="org-left">(min <sub>samples</sub> <sub>leaf</sub>=5)</td>
</tr></tbody></table> <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold; font-style: italic;">all_preds</span> = [t.predict(model_val_x)  <span style="font-weight: bold;">for</span> t  <span style="font-weight: bold;">in</span> trees]
 <span style="font-weight: bold; font-style: italic;">avg_preds</span> = np.stack(all_preds).mean(axis=0)
</pre>
</div>

 <div class="org-src-container">
 <pre class="src src-python">mean_absolute_error(model_val_y, avg_preds), balanced_logarithmic_loss(model_val_y.to_numpy(), avg_preds)
</pre>
</div>

 <p>
The competition metric gives the best result yet of 0.42. This is the
same score it received when submitted to the public leaderboard.
</p>
</div>
</div>
 <div id="outline-container-orgbd10eeb" class="outline-3">
 <h3 id="orgbd10eeb"> <span class="section-number-3">7.2.</span> Using  <code>sklearn</code>'s  <code>RandomForestClassifier</code></h3>
 <div class="outline-text-3" id="text-7-2">
 <p>
 <code>sklearn</code>'s  <code>RandomForestClassifier</code> does a similar process, but in addition
to selecting a random subset of rows, it also selects a random subset of
columns for each tree.
</p>

 <div class="org-src-container">
 <pre class="src src-python"> <span style="font-weight: bold;">from</span> sklearn.ensemble  <span style="font-weight: bold;">import</span> RandomForestClassifier

 <span style="font-weight: bold; font-style: italic;">rf</span> = RandomForestClassifier(n_estimators=100, min_samples_leaf=5)
rf.fit(model_train_x, model_train_y)
 <span style="font-weight: bold; font-style: italic;">preds</span> = rf.predict(model_val_x)

mean_absolute_error(model_val_y, preds), balanced_logarithmic_loss(model_val_y.to_numpy(), preds)
</pre>
</div>

 <table> <colgroup> <col class="org-right"></col> <col class="org-left"></col> <col class="org-left"></col></colgroup> <tbody> <tr> <td class="org-right">0.06451612903225806</td>
 <td class="org-left">np.float64</td>
 <td class="org-left">(5.318974763205966)</td>
</tr></tbody></table> <p>
Although locally the score seems to be worse by the competition
metric, on the public leaderboard it gave a score of 0.4, slightly
better than the handrolled random forest approach.
</p>
</div>
</div>
 <div id="outline-container-orgf1c9743" class="outline-3">
 <h3 id="orgf1c9743"> <span class="section-number-3">7.3.</span> Feature Importance</h3>
 <div class="outline-text-3" id="text-7-3">
 <p>
The random forest model can also tell us which features were most important
in making the predictions:
</p>

 <div class="org-src-container">
 <pre class="src src-python">pd.DataFrame( <span style="font-weight: bold;">dict</span>(cols=model_train_x.columns, imp=rf.feature_importances_)).plot(
     <span style="font-style: italic;">"cols"</span>,  <span style="font-style: italic;">"imp"</span>,  <span style="font-style: italic;">"barh"</span>, figsize=(8, 20)
)
</pre>
</div>


 <div id="orgdb78da6" class="figure">
 <p> <img src="img/icr_feature_importance.png" alt="icr_feature_importance.png"></img></p>
</div>

 <p>
From this, we can see that  <code>DU</code> is the column that most heavily
influences the final prediction. Since the data is anonymized, we
cannot tell what trait  <code>DU</code> represents, but it seems to noticeably
determine the outcome.
</p>
</div>
</div>
</div>
 <div id="outline-container-org6c5eecb" class="outline-2">
 <h2 id="org6c5eecb"> <span class="section-number-2">8.</span> Conclusion</h2>
 <div class="outline-text-2" id="text-8">
 <p>
In this notebook, we started with a very simple model, the
 <a href="https://link.springer.com/article/10.1023/A:1022631118932">OneR model</a>, that makes predictions based on just a single feature of the
dataset. This kind of classifier was actually found to perform competitively
with other machine learning methods of the early 90s.
</p>

 <p>
We improved on the OneR model by performing splits based on several
features instead of one, and thus implemented a decision tree. However, the
bigger decision trees are, the more they tend to overfit training data.
</p>

 <p>
Next, we saw how we could improve on a single decision tree by using many
trees together working in an ensemble, each making predictions and then
combining their results by averaging their predictions, which gave the best
score yet.
</p>

 <p>
However, there's still more that can be done to get better predictions
in this competition. The goal of this notebook was to demonstrate the
basic techniques that apply to tabular data.
</p>

 <p>
Inspired by Jeremy Howard's excellent notebook  <a href="https://www.kaggle.com/code/jhoward/how-random-forests-really-work/">here</a>.
</p>
</div>
</div>
</div>]]></content>
  <link href="https://dataai.ng/one_r_to_random_forests.html"/>
  <id>https://dataai.ng/one_r_to_random_forests.html</id>
  <updated>2026-07-28T14:04:00+03:00</updated>
</entry>
</feed>
