1<!DOCTYPE html> 2<html lang="en"><head> 3 <meta charset="utf-8"> 4 <meta http-equiv="X-UA-Compatible" content="IE=edge"> 5 <meta name="viewport" content="width=device-width, initial-scale=1"><!-- Begin Jekyll SEO tag v2.8.0 --> 6<title>Interpretable ECG Classification With 1D Vision Transformer | Yoni Gottesman</title> 7<meta name="generator" content="Jekyll v3.10.0" /> 8<meta property="og:title" content="Interpretable ECG Classification With 1D Vision Transformer" /> 9<meta name="author" content="Yoni Gottesman" /> 10<meta property="og:locale" content="en_US" /> 11<meta name="description" content="Interpretable ECG Classification With 1D Vision Transformer" /> 12<meta property="og:description" content="Interpretable ECG Classification With 1D Vision Transformer" /> 13<link rel="canonical" href="https://yonigottesman.github.io/ecg/vit/deep-learning/2023/01/20/ecg-vit.html" /> 14<meta property="og:url" content="https://yonigottesman.github.io/ecg/vit/deep-learning/2023/01/20/ecg-vit.html" /> 15<meta property="og:site_name" content="Yoni Gottesman" /> 16<meta property="og:type" content="article" /> 17<meta property="article:published_time" content="2023-01-20T04:48:38+00:00" /> 18<meta name="twitter:card" content="summary" /> 19<meta property="twitter:title" content="Interpretable ECG Classification With 1D Vision Transformer" />
20<script type="application/ld+json"> 21{"@context":"https://schema.org","@type":"BlogPosting","author":{"@type":"Person","name":"Yoni Gottesman"},"dateModified":"2023-01-20T04:48:38+00:00","datePublished":"2023-01-20T04:48:38+00:00","description":"Interpretable ECG Classification With 1D Vision Transformer","headline":"Interpretable ECG Classification With 1D Vision Transformer","mainEntityOfPage":{"@type":"WebPage","@id":"https://yonigottesman.github.io/ecg/vit/deep-learning/2023/01/20/ecg-vit.html"},"url":"https://yonigottesman.github.io/ecg/vit/deep-learning/2023/01/20/ecg-vit.html"}</script>
21 22<!-- End Jekyll SEO tag --> 23<link id="main-stylesheet" rel="stylesheet" href="/assets/css/style.css"><link type="application/atom+xml" rel="alternate" href="https://yonigottesman.github.io/feed.xml" title="Yoni Gottesman" />
vendor: 63 bytes, line 23
23<script async src="https://www.googletagmanager.com/gtag/js?id=
23G-H0P18BF8T9
vendor: 12 bytes, line 23
23"></script>
24<script> 25 window.dataLayer = window.dataLayer || []; 26 function gtag(){window.dataLayer.push(arguments);} 27 gtag('js', new Date()); 28 29 gtag('config', 'G-H0P18BF8T9'); 30</script>
30 31 32</head> 33<body><header class="site-header"> 34 35 <div class="wrapper"> 36 <a class="site-title" rel="author" href="/">Yoni Gottesman</a> 37 <nav class="site-nav"> 38 <input type="checkbox" id="nav-trigger"> 39 <label for="nav-trigger"> 40 <span class="menu-icon"></span> 41 </label> 42 43 <div class="nav-items"> 44 <a class="nav-item" href="/about/">About</a> 45</div> 46 47 </nav> 48 </div> 49</header> 50<main class="page-content" aria-label="Content"> 51 <div class="wrapper"> 52 <article class="post h-entry" itemscope itemtype="http://schema.org/BlogPosting"> 53 54 <header class="post-header"> 55 <h1 class="post-title p-name" itemprop="name headline">Interpretable ECG Classification With 1D Vision Transformer</h1> 56 <div class="post-meta"> 57 <time class="dt-published" datetime="2023-01-20T04:48:38+00:00" itemprop="datePublished"> 58 Jan 20, 2023 59 </time> 60 </div> 61 </header> 62 63 <div class="post-content e-content" itemprop="articleBody"> 64 <!-- Mathjax Support -->
65<script type="text/javascript" async="" src="https://cdn.mathjax.org/mathjax/latest/MathJax.js?config=TeX-MML-AM_CHTML"> 66</script>
66 67 68<p>In this post, I will use a vision transformer to classify ECG signals and use the attention scores to interpret what part of the signal the model is focusing on. 69All the code to reproduce the results is in <a href="https://github.com/yonigottesman/ecg_vit">my github</a>.</p> 70 71<h2 id="electrocardiogram-ecg">Electrocardiogram (ECG)</h2> 72<p>An ECG is a noninvasive test that records the heartâs electrical activity. This activity is the coordinated electrical impulses generated and transmitted throughout the heartâs muscle tissue, causing it to contract and pump blood. The test is performed by attaching electrodes to the skin of the chest, arms, and legs. These electrodes measure the voltage amplitude and direction of the heart activity over time.</p> 73 74<p>The following animation presents the heartâs electrical activity and how an ECG records it. The red dots/lines are an electrical pulse initiating at the right atrium, causing it to contract and pump deoxygenated blood to the right ventricle. At the same time, the left atrium also contracts and pumps oxygen-rich blood to the left ventricle. The electrical pulse then reaches the atrioventricular (AV) node between the atria and ventricles and is delayed before it is sent down into the ventricular muscle. This triggers the ventricles to contract and pump blood out to the body.<br> 75The blue recording at the bottom is the voltage recorded between two electrodes during the cardiac cycle.</p> 76 77<p><img src="https://upload.wikimedia.org/wikipedia/commons/0/0b/ECG_Principle_fast.gif" alt="" title="ecg-gif" height="50%" width="50%"> 78This is the recording of the voltage between only two electrodes, but to get a more detailed view of the heartâs activity, we attach more electrodes and measure the voltage from 12 different positions. In my mind, I imagen 12 cameras recording the heart from 12 3D locations.</p> 79 80<p>The following image illustrates the locations from which the 12 leads are recorded. The leadâs names are I, II, III, aVR, aVL, v1, v2, v3, v4, v5, and v6.<br> 81<img src="https://cdn.shopify.com/s/files/1/0059/3992/files/Image_5.png?v=1476239877" alt="" title="12-lead" height="50%" width="50%"></p> 82 83<p>The output of the ECG is the 12 leads recordings of voltage over time and looks like this (click to enlarge): 84<a href="/assets/ecgvit/ecg12.png"><img src="https://yonigottesman.github.io/assets/ecgvit/ecg12.png" alt=""></a> 85As you can see, every lead records a different view of the heartâs electrical activity. Also, you can see from the previous image that leads âaVRâ and âIâ are recording the activity from nearly opposite positions. Similarly, If you look at the recording of âaVRâ and âIâ you see that the signals are opposite.</p> 86 87<h4 id="hearbeat-complex">Hearbeat Complex</h4> 88<p>The cardiac cycle is the sequence of events that occur during one complete heartbeat. The different components of the cardiac cycle include:</p> 89<ul> 90 <li>The P-wave.</li> 91 <li>The QRS complex.</li> 92 <li>The T-wave.</li> 93 <li>The PR Interval.</li> 94 <li>The ST segment.</li> 95 <li>The QT interval.</li> 96</ul> 97 98<p>The following image displays each component: 99<img src="https://upload.wikimedia.org/wikipedia/commons/thumb/9/9e/SinusRhythmLabels.svg/1280px-SinusRhythmLabels.svg.png" alt="" title="pqrst" height="50%" width="50%"></p> 100 101<h2 id="left-ventricular-hypertrophy">Left Ventricular Hypertrophy</h2> 102<p>Physicians use the ECG to diagnose various diseases, and during this post, I will focus on one in particular: âLeft Ventricular Hypertrophyâ (LVH). LVH is a condition where (as the name suggests :) ) there is hypertrophy of the left ventricle heart muscle. Hypertrophy means a thickening of muscle tissue, and Left Ventricular Hypertrophy means a thickening of the muscle tissue of the left ventricle. 103The thickened heart tissue causes the ventricle lumen to be smaller and the left ventricular heart muscle to be less elastic, making it harder to pump blood to the rest of the body.</p> 104 105<p><img src="https://www.mayoclinic.org/-/media/kcms/gbs/patient-consumer/images/2013/08/26/11/03/ds00680-ans7_lvhthu_jpg.jpg" alt="" title="lvh" height="50%" width="50%"></p> 106 107<p>The manifestation of LVH in ECG exams is high voltage because the increased muscle mass of the left ventricle generates a larger electrical current, which results in a taller QRS complex. Another indication is inverted T waves. Here is an example of lead V5 from a patient with LVH (More on the dataset later). The R peaks are high, and the T waves are inverted:</p> 108 109<p><a href="/assets/ecgvit/lvhV5.png"><img src="https://yonigottesman.github.io/assets/ecgvit/lvhV5.png" alt=""></a></p> 110 111<p>After this short introduction to ECG, itâs time to train our model <img class="emoji" title=":robot:" alt=":robot:" src="https://github.githubassets.com/images/icons/emoji/unicode/1f916.png" height="20" width="20"></p> 112 113<h2 id="data">Data</h2> 114<p>Iâm going to use <a href="https://www.physionet.org/content/ecg-arrhythm
114ia/1.0.0/">this dataset</a> which contains 45152 12 lead ECGs. Each ECG comes with a list of labels described <a href="https://www.physionet.org/content/ecg-arrhythmia/1.0.0/ConditionNames_SNOMED-CT.csv">here</a>. The file contains descriptions of many conditions, but I am only interested in two: LVH and TWO (T-Wave Opposite).<br> 115The first step is to parse the files in the dataset and create a dataframe representing the training data. <a href="https://github.com/yonigottesman/ecg_vit/blob/main/data.py#L38">This function</a> parses all the files, looks for the LVH and TWO, labels, and builds the dataframe. To simplify, I take all the LVH/TWO ECGs and sample only 20000 negative ECGs. Finally, I split the data to train/test using a hash function on the filename. The resulting dataframe contains the <code class="language-plaintext highlighter-rouge">file</code> path, <code class="language-plaintext highlighter-rouge">y</code> label, which is the hot encoding of LVH and TWO, and <code class="language-plaintext highlighter-rouge">test</code>, which indicates validation data:</p> 116 117<table> 118 <thead> 119 <tr> 120 <th style="text-align: left">file</th> 121 <th style="text-align: left">y</th> 122 <th style="text-align: left">test</th> 123 </tr> 124 </thead> 125 <tbody> 126 <tr> 127 <td style="text-align: left"><root_path>/WFDBRecords/34/342/JS33623</td> 128 <td style="text-align: left">[0. 0.]</td> 129 <td style="text-align: left">True</td> 130 </tr> 131 <tr> 132 <td style="text-align: left"><root_path>/WFDBRecords/42/425/JS41970</td> 133 <td style="text-align: left">[0. 0.]</td> 134 <td style="text-align: left">False</td> 135 </tr> 136 <tr> 137 <td style="text-align: left"><root_path>/WFDBRecords/02/024/JS01566</td> 138 <td style="text-align: left">[0. 0.]</td> 139 <td style="text-align: left">False</td> 140 </tr> 141 <tr> 142 <td style="text-align: left"><root_path>/WFDBRecords/42/424/JS41856</td> 143 <td style="text-align: left">[0. 0.]</td> 144 <td style="text-align: left">False</td> 145 </tr> 146 <tr> 147 <td style="text-align: left"><root_path>/WFDBRecords/22/220/JS21462</td> 148 <td style="text-align: left">[1. 0.]</td> 149 <td style="text-align: left">False</td> 150 </tr> 151 <tr> 152 <td style="text-align: left"><root_path>/WFDBRecords/32/329/JS32343</td> 153 <td style="text-align: left">[0. 0.]</td> 154 <td style="text-align: left">False</td> 155 </tr> 156 </tbody> 157</table> 158 159<div class="alert alert-info" role="alert"> <b>Note:</b> Itâs better to split train/test by hash and not some random function to get the same split every time. Setting the random seed is not enough because if I add even one sample to the train, the whole train/test split will be different. 160</div> 161 162<p>I use the <code class="language-plaintext highlighter-rouge">wfdb</code> library to parse the files into an <code class="language-plaintext highlighter-rouge">np.array</code> representing the ECG. The <code class="language-plaintext highlighter-rouge">np.array</code> shape is (5000,12) - 5000 samples of 12 leads. In this dataset the sampling rate is 500, meaning the length of each ECG is 10 seconds.</p> 163<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">record</span> <span class="o">=</span> <span class="n">wfdb</span><span class="p">.</span><span class="n">rdrecord</span><span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="s">"file"</span><span class="p">].</span><span class="n">values</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span> 164<span class="k">print</span><span class="p">(</span><span class="n">record</span><span class="p">.</span><span class="n">p_signal</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span> 165<span class="k">print</span><span class="p">(</span><span class="n">record</span><span class="p">.</span><span class="n">fs</span><span class="p">)</span> 166 167<span class="o">></span> <span class="p">(</span><span class="mi">5000</span><span class="p">,</span> <span class="mi">12</span><span class="p">)</span> 168<span class="o">></span> <span class="mi">500</span> 169</code></pre></div></div> 170<p>Now I can create the TensorFlow train/val datasets. Nothing special here; shuffle, open files, batch, and prefetch.</p> 171 172<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">read_record</span><span class="p">(</span><span class="n">path</span><span class="p">):</span> 173 <span class="n">
173record</span> <span class="o">=</span> <span class="n">wfdb</span><span class="p">.</span><span class="n">rdrecord</span><span class="p">(</span><span class="n">path</span><span class="p">.</span><span class="n">decode</span><span class="p">(</span><span class="s">"utf-8"</span><span class="p">))</span> 174 <span class="k">return</span> <span class="n">record</span><span class="p">.</span><span class="n">p_signal</span><span class="p">.</span><span class="n">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span> 175 176<span class="k">def</span> <span class="nf">ds_base</span><span class="p">(</span><span class="n">df</span><span class="p">,</span> <span class="n">shuffle</span><span class="p">,</span> <span class="n">bs</span><span class="p">):</span> 177 <span class="n">ds</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">Dataset</span><span class="p">.</span><span class="n">from_tensor_slices</span><span class="p">((</span><span class="n">df</span><span class="p">[</span><span class="s">"file"</span><span class="p">],</span> <span class="nb">list</span><span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="s">"y"</span><span class="p">])))</span> 178 <span class="k">if</span> <span class="n">shuffle</span><span class="p">:</span> 179 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="n">shuffle</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">df</span><span class="p">))</span> 180 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="nb">map</span><span class="p">(</span> 181 <span class="k">lambda</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">:</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">numpy_function</span><span class="p">(</span><span class="n">read_record</span><span class="p">,</span> <span class="n">inp</span><span class="o">=</span><span class="p">[</span><span class="n">x</span><span class="p">],</span> <span class="n">Tout</span><span class="o">=</span><span class="n">tf</span><span class="p">.</span><span class="n">
181float32</span><span class="p">),</span> <span class="n">y</span><span class="p">),</span> 182 <span class="n">num_parallel_calls</span><span class="o">=</span><span class="n">tf</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">AUTOTUNE</span><span class="p">,</span> 183 <span class="n">deterministic</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> 184 <span class="p">)</span> 185 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="nb">map</span><span class="p">(</span><span class="k">lambda</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">:</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">where</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">math</span><span class="p">.</span><span class="n">is_nan</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">tf</span><span class="p">.</span><span class="n">zeros_like</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">x</span><span class="p">),</span> <span class="n">y</span><span class="p">))</span> 186 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="nb">map</span><span class="p">(</span><span class="k">lambda</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">:</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">ensure_shape</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="p">[</span><span class="mi">5000</span><span class="p">,</span> <span class="mi">12</span><span class="p">]),</span> <span class="n">y</span><span class="p">))</span> 187 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="n">batch</span><span class="p">(</span><span class="n">bs</span><span class="p">)</span> 188 <span class="n">ds</span> <span class="o">=</span> <span class="n">ds</span><span class="p">.</span><span class="n">prefetch</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">AUTOTUNE</span><span class="p">)</span> 189 <span class="k">return</span> <span class="n">ds</span> 190 191<span class="k">def</span> <span class="nf">gen_datasets</span><span class="p">(</span><span class="n">df</span><span class="p">,</span> <span class="n">bs</span><span class="p">):</span> 192 <span class="n">train_ds</span> <span class="o">=</span> <span class="n">ds_base</span><span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="o">~</span><span class="n">df</span><span class="p">[</span><span class="s">"test"</span><span class="p">]],</span> <span class="bp">True</span><span class="p">,</span> <span class="n">bs</span><span class="p">)</span> 193 <span class="n">val_ds</span> <span class="o">=</span> <span class="n">ds_base</span><span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="n">df</span><span class="p">[</span><span class="s">"test"</span><span class="p">]],</span> <span class="bp">False</span><span class="p">,</span> <span class="n">bs</span><span class="p">)</span> 194 <span class="k">return</span> <span class="n">train_ds</span><span class="p">,</span> <span class="n">val_ds</span> 195</code></pre></div></div> 196 197<h2 id="model---1d-vision-transformer">Model - 1D Vision Transformer</h2> 198<p>The model is a vision transformer with a minor change; the patches are not 2D 16x16, but 1D sized 20. Twenty samples represent 0.04[sec], precisely one small square on the ECG image. The small squares are the standard granularity used for diagnosis, so it makes sense for that to be the patch size. The Conv filter size is 20, but the number of channels in the signal is 12, so each patch is sized 20 across all leads (channels).</p> 199 200<p>Letâs start with the embeddings layer. The input is an <code class="language-plaintext highlighter-rouge">(5000,12)</code> ECG array, and the layer will create embeddings for <code class="language-plaintext highlighter-rouge">5000/patch_size+1</code> patches (including the <code class="language-plaintext highlighter-rouge">cls_token</code>). Then, add positional embeddings and return a tensor shaped <code class="language-plaintext highlighter-rouge">(bs,5000/patch_size+1,hidden_size)</code>.</p> 201<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">ViTEmbeddings</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">
201Layer</span><span class="p">):</span> 202 <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">patch_size</span><span class="p">,</span> <span class="n">hidden_size</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span> 203 <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">(</span><span class="o">**</span><span class="n">kwargs</span><span class="p">)</span> 204 205 <span class="bp">self</span><span class="p">.</span><span class="n">patch_size</span> <span class="o">=</span> <span class="n">patch_size</span> 206 <span class="bp">self</span><span class="p">.</span><span class="n">hidden_size</span> <span class="o">=</span> <span class="n">hidden_size</span> 207 208 <span class="bp">self</span><span class="p">.</span><span class="n">patch_embeddings</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Conv1D</span><span class="p">(</span><span class="n">filters</span><span class="o">=</span><span class="n">hidden_size</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="n">patch_size</span><span class="p">,</span> <span class="n">strides</span><span class="o">=</span><span class="n">patch_size</span><span class="p">)</span> 209 <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="n">rate</span><span class="o">=</span><span class="n">dropout</span><span class="p">)</span> 210 211 <span class="k">def</span> <span class="nf">build</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">input_shape</span><span class="p">):</span> 212 <span class="bp">self</span><span class="p">.</span><span class="n">cls_token</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">add_weight</span><span class="p">(</span><span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">hidden_size</span><span class="p">),</span> <span class="n">trainable</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="s">"cls_token"</span><span class="p">)</span> 213 214 <span class="n">num_patches</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">//</span> <span class="bp">self</span><span class="p">.</span><span class="n">patch_size</span> 215 <span class="bp">self</span><span class="p">.</span><span class="n">position_embeddings</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">add_weight</span><span class="p">(</span> 216 <span class="n">shape</span><span class="o">=</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">num_patches</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">hidden_size</span><span class="p">),</span> <span class="n">trainable</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="s">"position_embeddings"</span>
217 <span class="p">)</span> 218 219 <span class="k">def</span> <span class="nf">call</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">:</span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="nb">bool</span> <span class="o">=</span> <span class="bp">False</span><span class="p">)</span> <span class="o">-></span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span> 220 <span class="n">inputs_shape</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">shape</span><span class="p">(</span><span class="n">inputs</span><span class="p">)</span> <span class="c1"># N,H,W,C 221</span> <span class="n">embeddings</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">patch_embeddings</span><span class="p">(</span><span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 222 223 <span class="c1"># add the [CLS] token to the embedded patch tokens 224</span> <span class="n">cls_tokens</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">repeat</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">cls_token</span><span class="p">,</span> <span class="n">repeats</span><span class="o">=</span><span class="n">inputs_shape</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span> 225 <span class="n">embeddings</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">concat</span><span class="p">((</span><span class="n">cls_tokens</span><span class="p">,</span> <span class="n">embeddings</span><span class="p">),</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span> 226 227 <span class="c1"># add positional encoding to each token 228</span> <span class="n">embeddings</span> <span class="o">=</span> <span class="n">embeddings</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">position_embeddings</span> 229 <span class="n">embeddings</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">embeddings</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 230 231 <span class="k">return</span> <span class="n">embeddings</span> 232</code></pre></div></div> 233 234<p>Next is the MLP; nothing special here. It is the same as in the vit paper.</p> 235<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">MLP</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Layer</span><span class="p">):</span> 236 <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">mlp_dim</span><span class="p">,</span> <span class="n">out_dim</span><span class="o">=</span><span class="bp">None</span><span class="p">,</span> <span class="n">activation</span><span class="o">=</span><span class="s">"gelu"</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span> 237 <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">(</span><span class="o">
237**</span><span class="n">kwargs</span><span class="p">)</span> 238 <span class="bp">self</span><span class="p">.</span><span class="n">mlp_dim</span> <span class="o">=</span> <span class="n">mlp_dim</span> 239 <span class="bp">self</span><span class="p">.</span><span class="n">out_dim</span> <span class="o">=</span> <span class="n">out_dim</span> 240 <span class="bp">self</span><span class="p">.</span><span class="n">activation</span> <span class="o">=</span> <span class="n">activation</span> 241 <span class="bp">self</span><span class="p">.</span><span class="n">dropout_rate</span> <span class="o">=</span> <span class="n">dropout</span> 242 243 <span class="k">def</span> <span class="nf">build</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">input_shape</span><span class="p">):</span> 244 <span class="bp">self</span><span class="p">.</span><span class="n">dense1</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Dense</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">mlp_dim</span><span class="p">)</span> 245 <span class="bp">self</span><span class="p">.</span><span class="n">activation1</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Activation</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">activation</span><span class="p">)</span> 246 <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">dropout_rate</span><span class="p">)</span> 247 <span class="bp">self</span><span class="p">.</span><span class="n">dense2</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Dense</span><span class="p">(</span><span class="n">input_shape</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]</span> <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">out_dim</span> <span class="ow">is</span> <span class="bp">None</span> <span class="k">else</span> <span class="bp">self</span><span class="p">.</span><span class="n">out_dim</span><span class="p">)</span> 248 249 <span class="k">def</span> <span class="nf">call</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">:</span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="nb">bool</span> <span class="o">=</span> <span class="bp">False</span><span class="p">):</span> 250 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dense1</span><span class="p">(</span><span class="n">inputs</span><span class="p">)</span> 251 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">activation1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> 252 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span>
253 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dense2</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> 254 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">dropout</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 255 <span class="k">return</span> <span class="n">x</span> 256</code></pre></div></div> 257<p>The main encoder block. Normalization layers, attention layer, and MLP are all connected with skip connections with a stochastic depth layer. 258Each block has an additional function, <code class="language-plaintext highlighter-rouge">get_attention_scores</code>, that isnât called during training. <code class="language-plaintext highlighter-rouge">get_attention_scores</code> is used only to get the attention weights from the attention layer, and I will use it later for interpretability.</p> 259<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">Block</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Layer</span><span class="p">):</span> 260 <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span> 261 <span class="bp">self</span><span class="p">,</span> 262 <span class="n">num_heads</span><span class="p">,</span> 263 <span class="n">attention_dim</span><span class="p">,</span> 264 <span class="n">attention_bias</span><span class="p">,</span> 265 <span class="n">mlp_dim</span><span class="p">,</span> 266 <span class="n">attention_dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> 267 <span class="n">sd_survival_probability</span><span class="o">=</span><span class="mf">1.0</span><span class="p">,</span> 268 <span class="n">activation</span><span class="o">=</span><span class="s">"gelu"</span><span class="p">,</span> 269 <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> 270 <span class="o">**</span><span class="n">kwargs</span><span class="p">,</span> 271 <span class="p">):</span> 272 <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">(</span><span class="o">**</span><span class="n">kwargs</span><span class="p">)</span> 273 <span class="bp">self</span><span class="p">.</span><span class="n">norm_before</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">LayerNormalization</span><span class="p">()</span> 274 <span class="bp">self</span><span class="p">.</span><span class="n">attn</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">MultiHeadAttention</span><span class="p">(</span> 275 <span class="n">num_heads</span><span class="p">,</span> 276 <span class="n">attention_dim</span> <span class="o">//</span> <span class="n">num_heads</span><span class="p">,</span> 277 <span class="n">use_bias</span><span class="o">=</span><span class="n">attention_bias</span><span class="p">,</span> 278 <span class="n">dropout</span><span class="o">=</span><span class="n">attention_dropout</span><span class="p">,</span> 279 <span class="p">)</span> 280 <span class="bp">self</span><span class="p">.</span><span class="n">stochastic_depth</span> <span class="o">=</span> <span class="n">tfa</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">StochasticDepth</span><span class="p">(</span><span class="n">sd_survival_probability</span><span class="p">)</span>
281 <span class="bp">self</span><span class="p">.</span><span class="n">norm_after</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">LayerNormalization</span><span class="p">()</span> 282 <span class="bp">self</span><span class="p">.</span><span class="n">mlp</span> <span class="o">=</span> <span class="n">MLP</span><span class="p">(</span><span class="n">mlp_dim</span><span class="o">=</span><span class="n">mlp_dim</span><span class="p">,</span> <span class="n">activation</span><span class="o">=</span><span class="n">activation</span><span class="p">,</span> <span class="n">dropout</span><span class="o">=</span><span class="n">dropout</span><span class="p">)</span> 283 284 <span class="k">def</span> <span class="nf">build</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">input_shape</span><span class="p">):</span> 285 <span class="nb">super</span><span class="p">().</span><span class="n">build</span><span class="p">(</span><span class="n">input_shape</span><span class="p">)</span> 286 <span class="c1"># TODO YONIGO: tf doc says to do this ¯\_(ã)_/¯ 287</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn</span><span class="p">.</span><span class="n">_build_from_signature</span><span class="p">(</span><span class="n">input_shape</span><span class="p">,</span> <span class="n">input_shape</span><span class="p">)</span> 288 289 <span class="k">def</span> <span class="nf">call</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="bp">False</span><span class="p">):</span> 290 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">norm_before</span><span class="p">(</span><span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 291 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 292 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">stochastic_depth</span><span class="p">([</span><span class="n">inputs</span><span class="p">,</span> <span class="n">x</span><span class="p">],</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 293 <span class="n">x2</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">norm_after</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 294 <span class="n">x2</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mlp</span><span class="p">(</span><span class="n">x2</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 295 <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">stochastic_depth</span><span class="p">([</span><span class="n">x</span><span class="p">,</span> <span class="n">x2</span><span class="p">],</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 296 297 <span class="k">def</span> <span class="nf">get_attention_scores</span><span class="p">(</span><span class="bp">
297self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">):</span> 298 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">norm_before</span><span class="p">(</span><span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span> 299 <span class="n">_</span><span class="p">,</span> <span class="n">weights</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">attn</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">return_attention_scores</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> 300 <span class="k">return</span> <span class="n">weights</span> 301</code></pre></div></div> 302<p>Finally, I tie everything together in the VisionTransformer model:</p> 303<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">VisionTransformer</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">Model</span><span class="p">):</span> 304 <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span> 305 <span class="bp">self</span><span class="p">,</span> 306 <span class="n">patch_size</span><span class="p">,</span> 307 <span class="n">hidden_size</span><span class="p">,</span> 308 <span class="n">depth</span><span class="p">,</span> 309 <span class="n">num_heads</span><span class="p">,</span> 310 <span class="n">mlp_dim</span><span class="p">,</span> 311 <span class="n">num_classes</span><span class="p">,</span> 312 <span class="n">dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> 313 <span class="n">sd_survival_probability</span><span class="o">=</span><span class="mf">1.0</span><span class="p">,</span> 314 <span class="n">attention_bias</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> 315 <span class="n">attention_dropout</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span> 316 <span class="o">*</span><span class="n">args</span><span class="p">,</span> 317 <span class="o">**</span><span class="n">kwargs</span><span class="p">,</span> 318 <span class="p">):</span> 319 <span class="nb">super</span><span class="p">().</span><span class="n">__init__</span><span class="p">(</span><span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span> 320 321 <span class="bp">self</span><span class="p">.</span><span class="n">embeddings</span> <span class="o">=</span> <span class="n">ViTEmbeddings</span><span class="p">(</span><span class="n">patch_size</span><span class="p">,</span> <span class="n">hidden_size</span><span class="p">,</span> <span class="n">dropout</span><span class="p">)</span> 322 <span class="n">sd</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">linspace</span><span class="p">(</span><span class="mf">1.0</span><span class="p">,</span> <span class="n">sd_survival_probability</span><span class="p">,</span> <span class="n">depth</span><span class="p">)</span> 323 <span class="bp">self</span><span class="p">.</span><span class="n">blocks</span> <span class="o">=</span> <span class="p">[</span> 324 <span class="n">Block</span><span class="p">(</span> 325 <span class="n">num_heads</span><span class="p">,</span> 326 <span class="n">attention_dim</span><span class="o">=</span><span class="n">hidden_size</span><span class="p">,</span>
327 <span class="n">attention_bias</span><span class="o">=</span><span class="n">attention_bias</span><span class="p">,</span> 328 <span class="n">attention_dropout</span><span class="o">=</span><span class="n">attention_dropout</span><span class="p">,</span> 329 <span class="n">mlp_dim</span><span class="o">=</span><span class="n">mlp_dim</span><span class="p">,</span> 330 <span class="n">sd_survival_probability</span><span class="o">=</span><span class="p">(</span><span class="n">sd</span><span class="p">[</span><span class="n">i</span><span class="p">].</span><span class="n">numpy</span><span class="p">().</span><span class="n">item</span><span class="p">()),</span> 331 <span class="n">dropout</span><span class="o">=</span><span class="n">dropout</span><span class="p">,</span> 332 <span class="p">)</span> 333 <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">depth</span><span class="p">)</span> 334 <span class="p">]</span> 335 336 <span class="bp">self</span><span class="p">.</span><span class="n">norm</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">LayerNormalization</span><span class="p">()</span> 337 338 <span class="bp">self</span><span class="p">.</span><span class="n">head</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">layers</span><span class="p">.</span><span class="n">Dense</span><span class="p">(</span><span class="n">num_classes</span><span class="p">)</span> 339 340 <span class="k">def</span> <span class="nf">call</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">:</span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="nb">bool</span> <span class="o">=</span> <span class="bp">False</span><span class="p">)</span> <span class="o">-></span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">:</span> 341 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">embeddings</span><span class="p">(</span><span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 342 <span class="k">for</span> <span class="n">block</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">blocks</span><span class="p">:</span> 343 <span class="n">x</span> <span class="o">=</span> <span class="n">block</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="n">training</span><span class="p">)</span> 344 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">norm</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> 345 <span class="n">x</span> <span class="o">=</span> <span class="n">x</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">]</span> <span class="c1"># take only cls_token 346</span> <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">head</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> 347 348 <span class="k">def</span> <span class="nf">get_last_selfattention</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">inputs</span><span class="p">:</span> <span class="n">tf</span><span class="p">.</span><span class="n">Tensor</span><span class="p">):</span>
349 <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">embeddings</span><span class="p">(</span><span class="n">inputs</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span> 350 <span class="k">for</span> <span class="n">block</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">blocks</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">]:</span> 351 <span class="n">x</span> <span class="o">=</span> <span class="n">block</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">training</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span> 352 <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">blocks</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">].</span><span class="n">get_attention_scores</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> 353</code></pre></div></div> 354 355<h2 id="train">Train</h2> 356<p>Create ViT and train!</p> 357<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">vit</span> <span class="o">=</span> <span class="n">VisionTransformer</span><span class="p">(</span> 358 <span class="n">patch_size</span><span class="o">=</span><span class="mi">20</span><span class="p">,</span> 359 <span class="n">hidden_size</span><span class="o">=</span><span class="mi">768</span><span class="p">,</span> 360 <span class="n">depth</span><span class="o">=</span><span class="mi">6</span><span class="p">,</span> 361 <span class="n">num_heads</span><span class="o">=</span><span class="mi">6</span><span class="p">,</span> 362 <span class="n">mlp_dim</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> 363 <span class="n">num_classes</span><span class="o">=</span><span class="nb">len</span><span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="s">"y"</span><span class="p">].</span><span class="n">values</span><span class="p">[</span><span class="mi">0</span><span class="p">]),</span> 364 <span class="n">sd_survival_probability</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> 365<span class="p">)</span> 366 367<span class="n">optimizer</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">optimizers</span><span class="p">.</span><span class="n">Adam</span><span class="p">(</span><span class="mf">0.0001</span><span class="p">)</span> 368<span class="n">loss</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">losses</span><span class="p">.</span><span class="n">BinaryCrossentropy</span><span class="p">(</span><span class="n">from_logits</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> 369<span class="n">metrics</span> <span class="o">=</span> <span class="p">[</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">metrics</span><span class="p">.</span><span class="n">AUC</span><span class="p">(</span><span class="n">from_logits</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="s">"roc_auc"</span><span class="p">)]</span> 370<span class="n">vit</span><span class="p">.</span><span class="nb">compile</span><span class="p">(</span><span class="n">optimizer</span><span class="o">=</span><span class="n">optimizer</span><span class="p">,</span> <span class="n">loss</span><span class="o">=</span><span class="n">loss</span><span class="p">,</span> <span class="n">metrics</span><span class="o">=</span><span class="n">metrics</span><span class="p">)</span> 371 372<span class="n">cbs</span> <span class="o">=</span> <span class="p">[</span><span class="n">tf</span><span class="p">.</span><span class="n">keras</span><span class="p">.</span><span class="n">callbacks</span><span class="p">.</span><span class="n">ModelCheckpoint</span><span class="p">(</span><span class="s">"vit_best/"</span><span class="p">,</span> <span class="n">
372monitor</span><span class="o">=</span><span class="s">"val_roc_auc"</span><span class="p">,</span> <span class="n">save_best_only</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">save_weights_only</span><span class="o">=</span><span class="bp">True</span><span class="p">)]</span> 373 374<span class="n">vit</span><span class="p">.</span><span class="n">fit</span><span class="p">(</span><span class="n">train_ds</span><span class="p">,</span> <span class="n">validation_data</span><span class="o">=</span><span class="n">val_ds</span><span class="p">,</span> <span class="n">epochs</span><span class="o">=</span><span class="mi">10</span><span class="p">,</span> <span class="n">callbacks</span><span class="o">=</span><span class="n">cbs</span><span class="p">)</span> 375</code></pre></div></div> 376 377<p>The average AUC of the roc curve on the validation set is <code class="language-plaintext highlighter-rouge">0.95</code>. This performance metric means nothing in the context of this blog post except that it is much better than random, and the model probably learned something. Is this result good or bad depends on other factors such as product definition, human-level performance on the task, and other benchmarks on the same or related tasks. I can improve performance by using more data, labels, augmentations, and signal preprocessing, such as cleaning high-frequency noise and low-frequency trends. But, the goal of this experiment is not to achieve SOTA using ViT but to see if I can use ViT for interpretability.</p> 378 379<h2 id="interpetability">Interpetability</h2> 380<p>Finally, the main course of this post! Did the model learn something meaningful? is it focusing on the same parts of the ECG as physicians do to detect LVH??? Remember, physicians look for high QRS complexes and inverted T-waves. 381I added to the ViT implementation a function <code class="language-plaintext highlighter-rouge">get_last_selfattention</code>. This function does a complete forward pass on all the layers except for the last block, on which it just calls <code class="language-plaintext highlighter-rouge">get_attention_scores</code> to return the attention scores from the <code class="language-plaintext highlighter-rouge">MultiHeadAttention</code> layer.</p> 382 383<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">record</span> <span class="o">=</span> <span class="n">wfdb</span><span class="p">.</span><span class="n">rdrecord</span><span class="p">(</span><span class="n">file_path</span><span class="p">)</span> 384<span class="n">attn</span> <span class="o">=</span> <span class="n">vit</span><span class="p">.</span><span class="n">get_last_selfattention</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">expand_dims</span><span class="p">(</span><span class="n">record</span><span class="p">.</span><span class="n">p_signal</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> 385<span class="k">print</span><span class="p">(</span><span class="n">attn</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span> 386<span class="o">></span> <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">6</span><span class="p">,</span> <span class="mi">251</span><span class="p">,</span> <span class="mi">251</span><span class="p">)</span> 387</code></pre></div></div> 388<p>The <code class="language-plaintext highlighter-rouge">attn</code> scores is shaped (<code class="language-plaintext highlighter-rouge">batch_size</code>, <code class="language-plaintext highlighter-rouge">num_heads</code>, <code class="language-plaintext highlighter-rouge">num_patches</code>, <code class="language-plaintext highlighter-rouge">num_patches</code>). The fully connected classification head at the end of the model is feeded only by the <code class="language-plaintext highlighter-rouge">cls_token</code> embedding, so Iâm interested in the attention scores of all the other <code class="language-plaintext highlighter-rouge">250</code> patches when the last attention layer calculated the <code class="language-plaintext highlighter-rouge">cls_token</code> output. 389The <code class="language-plaintext highlighter-rouge">cls_token</code> embedding is the first from the <code class="language-plaintext highlighter-rouge">251</code> (thatâs just how I did it in the <code class="language-plaintext highlighter-rouge">VitEmbeddings</code> layer), so for each of the <code class="language-plaintext highlighter-rouge">6</code> heads, I take row <code class="language-plaintext highlighter-rouge">0</code> and skip column <code class="language-plaintext highlighter-rouge">0</code> because Iâm not interested in the score of the <code class="language-plaintext highlighter-rouge">cls_token</code> on itself:</p> 390<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">attn</span> <span class="o">=</span> <span class="n">attn</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="p">:,</span> <span class="mi">0</span><span class="p">,</span> <span class="mi">1</span><span class="p">:]</span> 391<span class="k">print</span><span class="p">(</span><span class="n">attn</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span> 392<span class="o">></span> <span class="p">(</span><span class="mi">6</span><span class="p">,</span> <span class="mi">250</span><span class="p">)</span> 393</code></pre></div></div> 394<p>Now the <code class="language-plaintext highlighter-rouge">attn</code> contains the score of all <code class="language-plaintext highlighter-rouge">250</code> patches for each head. To display these scores on an ECG, I need to resize the <code class="language-plaintext highlighter-rouge">250</code> back to <code class="language-plaintext highlighter-rouge">5000</code>. Each patch represents <code class="language-plaintext highlighter-rouge">20</code> samples from the ECG, so I need to repeat each score <code class="language-plaintext highlighter-rouge">20</code> times.</p> 395<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">attn</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">transpose</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> 396<span class="n">attn</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">expand_dims</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">expand_dims</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="mi">0</span><span class="p">)</span> 397<span class="n">attn</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">image</span><span class="p">.</span><span class="n">resize</span><span class="p">(</span><span class="n">attn</span><span class="p">,</span> <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">5000</span><span class="p">
397))[</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">]</span> 398<span class="k">print</span><span class="p">(</span><span class="n">attn</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span> 399<span class="o">></span> <span class="p">(</span><span class="mi">5000</span><span class="p">,</span> <span class="mi">6</span><span class="p">)</span> 400</code></pre></div></div> 401 402<p>Thatâs it! I can now plot an ECG lead and display the scores on top of it. In pyplot, I use <code class="language-plaintext highlighter-rouge">cmap=Reds</code> so that it colors high scores in darker red, and here are plots of scores from different attention heads:<br> 403<a href="/assets/ecgvit/attn_4.png"><img src="https://yonigottesman.github.io/assets/ecgvit/attn_4.png" alt=""></a> 404<a href="/assets/ecgvit/attn_2.png"><img src="https://yonigottesman.github.io/assets/ecgvit/attn_2.png" alt=""></a></p> 405 406<p>The first head learned to pay attention to the T-Wave, and the second to the QRS complex! The rest of the heads were a mix of QRS and T, so I didnât include them. 407Itâs also interesting to see the second head paying the most attention to the two last heartbeats. There is redundancy in these signals as the heartbeats donât vary much, so maybe the model learns to focus only on some of them.</p> 408 409<hr> 410<p><br></p> 411 412<p>Thats it! Not only can I use the model for ECG classification, but I can also explain what parts of the ECG the model pays attention to.</p> 413
414<script src="https://utteranc.es/client.js" repo="yonigottesman/yonigottesman.github.io" issue-term="pathname" label="comment" theme="github-light" crossorigin="anonymous" async=""> 415</script>
415 416 417 418 </div> 419 420 <a class="u-url" href="/ecg/vit/deep-learning/2023/01/20/ecg-vit.html" hidden></a> 421</article> 422 423 </div> 424 </main><link id="fa-stylesheet" rel="stylesheet" href="https://cdn.jsdelivr.net/npm/@fortawesome/[email protected]/css/all.min.css"> 425 426<footer class="site-footer h-card"> 427 <data class="u-url" value="/"></data> 428 429 <div class="wrapper"> 430 431 <div class="footer-col-wrapper"> 432 <div class="footer-col"> 433 <ul class="contact-list"> 434 <li class="p-name">Yoni Gottesman</li> 435 <li><a class="u-email" href="mailto:[email protected]">[email protected]</a></li> 436 </ul> 437 </div> 438 <div class="footer-col"> 439 <p>Random stuff I've learned</p> 440 </div> 441 </div> 442 443 <div class="social-links"> 444<ul class="social-media-list"> 445<li> 446 <a rel="me" href="https://github.com/yonigottesman" target="_blank" title="GitHub"> 447 <span class="grey fa-brands fa-github fa-lg"></span> 448 </a> 449 </li> 450 <li> 451 <a href="https://yonigottesman.github.io/feed.xml" target="_blank" title="Subscribe to syndication feed"> 452 <svg class="svg-icon grey" viewbox="0 0 16 16"> 453 <path d="M12.8 16C12.8 8.978 7.022 3.2 0 3.2V0c8.777 0 16 7.223 16 16h-3.2zM2.194 454 11.61c1.21 0 2.195.985 2.195 2.196 0 1.21-.99 2.194-2.2 2.194C.98 16 0 15.017 0 455 13.806c0-1.21.983-2.195 2.194-2.195zM10.606 456 16h-3.11c0-4.113-3.383-7.497-7.496-7.497v-3.11c5.818 0 10.606 4.79 10.606 10.607z"></path> 457 </svg> 458 </a> 459 </li> 460</ul> 461</div> 462 463 </div> 464 465</footer> 466 467</body> 468 469</html>
Line numbers count LF bytes from the start of the resource, as the search results do. Vendor segments are library code the classifier recognised; they are stored but not indexed. Bytes are shown as Latin1 characters, one per byte.