PageSourceSearch

https://yonigottesman.github.io/ecg/vit/deep-learning/2023/01/20/ecg-vit.html

html yonigottesman.github.io collected 2026-10-03 10:13:17 UTC 65,035 bytes, 469 lines download raw bytes

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">&lt;root_path&gt;/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">&lt;root_path&gt;/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">&lt;root_path&gt;/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">&lt;root_path&gt;/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">&lt;root_path&gt;/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">&lt;root_path&gt;/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">&gt;</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">&gt;</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">-&gt;</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">-&gt;</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">&gt;</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">&gt;</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">&gt;</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.