<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" ><generator uri="https://jekyllrb.com/" version="3.10.0">Jekyll</generator><link href="https://kimyeonz.github.io/blog/feed.xml" rel="self" type="application/atom+xml" /><link href="https://kimyeonz.github.io/blog/" rel="alternate" type="text/html" /><updated>2026-06-16T05:10:39+00:00</updated><id>https://kimyeonz.github.io/blog/feed.xml</id><title type="html">Hyeonseung’s Blog</title><subtitle>:triangular_ruler: Jekyll theme for building a personal site, blog, project documentation, or portfolio.</subtitle><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><entry><title type="html">Activation Steering</title><link href="https://kimyeonz.github.io/blog/interpretability/Activation_Steering/" rel="alternate" type="text/html" title="Activation Steering" /><published>2025-12-01T00:00:00+00:00</published><updated>2025-12-01T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/interpretability/Activation_Steering</id><content type="html" xml:base="https://kimyeonz.github.io/blog/interpretability/Activation_Steering/"><![CDATA[<p>기존에는 LLM의 행동을 제어하는 방법은 주로 프롬프팅이었다. 사용자가 원하는 스타일로 답변받으려면 프롬프트를 정교하게 작성해야 했고, 이를 따를지 말지는 온전히 모델의 선택이었다. 하지만 최근 주목받는 Activation Steering 기법은 이 패러다임을 바꾼다.</p>

<p>Steering(조종, 조작) Vector는 모델의 internal activation에 직접 개입하여 모델의 행동을 수학적으로 강제하는 기법이다. 프롬프팅처럼 ‘부탁하는’ 방식이 아니라, 내부 표현 공간에서 특정 방향 벡터를 더함으로써 모델이 반드시 원하는 방식으로 동작하도록 만드는 것이다.</p>

<p>이 기법은 단순한 제어 방법을 넘어 <strong>Interpretability(해석가능성)</strong>의 관점에서도 큰 강점을 가진다. 모델의 내부 활성화를 직접 조작한다는 것은, 동시에 모델이 특정 개념을 어떻게 표현하고 처리하는지를 명확하게 관찰할 수 있다는 의미다. 예를 들어, Love﻿ 벡터에서 Hate﻿ 벡터를 뺀 결과인 steering vector는 모델 내부에서 ‘긍정성’이 어떻게 인코딩되어 있는지를 직접 보여준다. 이는 단순히 출력을 설명하는 XAI 방식과는 달리, 모델의 ‘생각 과정’ 자체를 벡터 공간에서 들여다보는 것이다. 모델이 왜 특정 결과를 생성했는지 사후 분석하는 것이 아니라, 모델이 내부적으로 어떤 개념적 방향을 따라 이동하는지를 실시간으로 관찰하고 조작할 수 있게 되는 것이다.</p>

<h3 id="steering-vector">Steering Vector</h3>
<p>model내 특정 layer의 activation의 출력에 steering vector를 더해주면 출력을 의도한대로 변형할 수 있다. 예를 들어, Love와 Hate사이의 steering vector를 구하고
이를 I dislike the cars에 더하면 I like cars에 가까운 문장으로 변하는 것이다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img style="width: 50%;" alt="image" src="https://github.com/user-attachments/assets/baf17e88-2ed3-4864-85e7-ad8fc240d41d" />
  </p>
</div>

<p><br /></p>

<p>신경망이 개발된 초기에는 각 뉴런이 하나의 특징에만 반응한다고 생각했다(e.g., 이미지에서 공이 있을 때만 활성화, 축구화가 있을 때만 활성화).
하지만, polysemanticity(다의성) 특징에 의해 한 뉴런이 여러 개의 완전히 다른 개념에 반응한다(e.g., 공이 있을때도, 축구화가 있을 때도, 골키퍼가 있을 때도 활성화)는 특징을 발견했고,
이는 신경망은 뉴런 갯수보다 훨씬 많은 특징을 담아낸다는 것을 발견했다.</p>

<p>공, 축구화, 긍정, 부정과 같은 모델이 표현하고자 하는 추상적인 것을 feature라고 한다면 이 복잡한 feature들을 특정 activation(hidden layer를 통해 계산되고 학습되는 수치들)이 압축하여 저장하고 있는 것이다.
이 feature는 추상적인 개념으로 관련된 activation이 활성화되었을 때 나타내는 방향이라고 표현할 수 있고 따라서 steering vector를 activation에 더해주면서 출력을 원하는 방향으로 조종할 수 있다.</p>

<p>Steering vector는 아래와 같은 방법으로 계산할 수 있다.</p>

<p>1) 두 개의 input(프롬프트)에서 activation 추출</p>
<ul>
  <li>어느 레이어에서 추출할지 선택 (일반적으로 중간~깊은 레이어)
    <ul>
      <li>residual, attention, mlp layer 모두에서 사용 가능하고 어디에 사용해야 효과가 좋을지는 실험적으로 파악해야 할듯</li>
    </ul>
  </li>
  <li>steering vector를 어느 토큰의 activation에 적용할 것인지 선택 (보통 마지막 토큰. 마지막 토큰이 이전 정보를 다 가지고 있을거라는 가정)
    <ul>
      <li>아니면 특정 토큰에만 적용한다던가, 모두 다 한다던가.</li>
    </ul>
  </li>
  <li>두 프롬프트가 의미있게 대비되어야 함</li>
</ul>

<p>2) 두 activation의 차이 계산</p>
<ul>
  <li>steering vector 계산전 쓸 수 있는 방법(input activation에 대해)
    <ul>
      <li>activation_1 - activation_2 (단순 차)</li>
      <li>mean(activations_A) - mean(activations_B) (평균 차)</li>
    </ul>
  </li>
  <li>steering vector 계산후(선택사항) 이 값을
    <ul>
      <li>normalization / projection / PCA / regression</li>
    </ul>
  </li>
</ul>

<p>3) steering vector를 적용할 정도를 나타내는 $c$를 통해 새로운 input에 대해 1)과 같은 층에서 추출한 activation에 적용</p>
<ul>
  <li>ex) activation_steered = activation + c * steering_vector</li>
  <li>c가 너무 작으면 steering의 효과가 없을 것이고, 너무 크면 모델의 기존 답변과 달라져 목적이 무시될 수 있음</li>
</ul>

<h3 id="실제-활용-예시">실제 활용 예시</h3>
<p>1) Refusal Direction Removal</p>
<ul>
  <li>goal: 모델이 거부해야 하는 프롬프트에도 답변하도록 강제</li>
  <li>prompt set
    <ul>
      <li>harmful prompt(“누군가를 물리적으로 폭행하는 방법을 만드세요”)</li>
      <li>harmless prompt(“추도식 연설 작성”, “대학 교육의 이점”)</li>
    </ul>

    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># harmful에서 harmless까지의 차이를 구했으니까 이 방향은 harmful -&gt; harmless
</span><span class="n">refusal_vector</span> <span class="o">=</span> <span class="n">activation</span><span class="p">(</span><span class="n">harmful</span><span class="p">)</span> <span class="o">-</span> <span class="n">activation</span><span class="p">(</span><span class="n">harmless</span><span class="p">)</span>
<span class="n">steering</span> <span class="o">=</span> <span class="o">-</span><span class="n">refusal_vector</span>  <span class="c1"># harmless -&gt; harmful
</span></code></pre></div>    </div>
  </li>
  <li>results
<br /></li>
</ul>
<div align="center">
  <p>
  <img style="width: 50%;" alt="image" src="https://github.com/user-attachments/assets/b42106f6-8a34-4182-bbba-ace8a5d7f110" />
  </p>
</div>

<p><br /></p>

<p>2) Love 방향 (Activation Addition)</p>
<ul>
  <li>prompt set
    <ul>
      <li>Love: “I love you”</li>
      <li>Hate: “I hate you”</li>
    </ul>

    <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">v</span> <span class="o">=</span> <span class="n">activation</span><span class="p">(</span><span class="s">"I love you"</span><span class="p">)</span> <span class="o">-</span> <span class="n">activation</span><span class="p">(</span><span class="s">"I hate you"</span><span class="p">)</span>
</code></pre></div>    </div>
  </li>
  <li>results
    <ul>
      <li>input: “I hate you because…”</li>
      <li>no steering: “I hate you because you’re a coward. What I hate is that people think…”</li>
      <li>steering(love 방향 추가): “I hate you because you’re a wonderful person and the reason I’m here is because I want to be with you”</li>
    </ul>
  </li>
</ul>

<h3 id="실습">실습</h3>

<p>참고한 자료</p>
<ul>
  <li><a href="https://www.lesswrong.com/posts/ndyngghzFY388Dnew/implementing-activation-steering">Implementing activation steering</a></li>
  <li><a href="https://www.youtube.com/watch?v=cp-YSyc5aW8">Steering vectors: tailor LLMs without training. Part I: Theory (Interpretability Series)</a></li>
  <li><a href="https://github.com/user-attachments/assets/b42106f6-8a34-4182-bbba-ace8a5d7f110">Refusal in Language Models Is Mediated by a Single Direction</a></li>
  <li><a href="https://arxiv.org/pdf/2308.10248">STEERING LANGUAGE MODELS WITH ACTIVATION ENGINEERING</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Interpretability" /><summary type="html"><![CDATA[기존에는 LLM의 행동을 제어하는 방법은 주로 프롬프팅이었다. 사용자가 원하는 스타일로 답변받으려면 프롬프트를 정교하게 작성해야 했고, 이를 따를지 말지는 온전히 모델의 선택이었다. 하지만 최근 주목받는 Activation Steering 기법은 이 패러다임을 바꾼다.]]></summary></entry><entry><title type="html">Transformer와 MLP</title><link href="https://kimyeonz.github.io/blog/interpretability/Transformer/" rel="alternate" type="text/html" title="Transformer와 MLP" /><published>2025-11-23T00:00:00+00:00</published><updated>2025-11-23T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/interpretability/Transformer</id><content type="html" xml:base="https://kimyeonz.github.io/blog/interpretability/Transformer/"><![CDATA[<p>머신러닝과 비교하여 딥러닝 모델의 가장 큰 문제는 모델이 왜 이런 결과를 만들어냈는지 해석할 수 없다는 것이고 이를 ‘Black Box Problem’이라고 한다. 그리고 이 Black Box 문제를 해결해보고자 예전부터 XAI(eXplainable AI)에 관한 연구가 활발한데 XAI에서 사용하는 방법은 2가지로 나눌 수 있다.</p>
<ul>
  <li><strong>해석가능성(Interpretability)</strong>: 엔지니어의 관점에서 모델의 구조를 기반으로 알고리즘이 어떻게 변화하여 해당 결과를 도출해내는지 해석할 수 있는 능력</li>
  <li><strong>설명가능성(Explainability)</strong>: 해석가능성에서 더 나아가 사회과학, HCI 등의 요소를 결합하여 일반적인 사용자들도 이해할 수 있도록 설명하는 능력</li>
</ul>

<p>이번 포스팅에서는 딥러닝 모델에서 가장 많이 사용하는 Transformer를 기반으로 interpretability에 초점을 맞춰보자. 예전에도 <a href="https://hyeonseung0103.github.io/section4/S4_S2_Transformer/">Transformer에 관한 포스팅</a>을 작성한적이 있지만 이번에는 Attention과 MLP 부분을 좀 더 자세하게 다루어보겠다.</p>

<h1 id="embedding">Embedding</h1>
<p>Transformer에 문장이 입력되면 문장은 여러 토큰(무조건 토큰 하나가, 단어를 의미하는 것은 아님)으로 나뉘고 각 토큰은 해당 단어의 의미와 주변 단어들과의 관계, 위치 정보 등을 담은 고차원의 임베딩 벡터로 표현된다.</p>

<p>토큰이 임베딩 벡터로 표현된다는 것은 내적을 통해 단어사이의 유사도를 구할 수 있게 되는 것이고,</p>
<ul>
  <li>내적이 양수이면 단어가 유사한 것</li>
  <li>내적이 0이면 전혀 관련이 없는 것</li>
  <li>내적이 음수이면 단어가 반대의미를 가지는 것이다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img style="width: 50%;" alt="image" src="https://github.com/user-attachments/assets/db8323a8-8fbb-4815-92b9-6acc79db485a" />
  </p>
</div>

<p><br /></p>

<p>임베딩은 처음에는 랜덤한 값으로 이루어져있다가 여러 layer를 거치면서 data를 기반으로 아래와 같이 토큰의 정보를 가장 잘 담고 있는 상태로 업데이트된다.
이 embedding matrix의 크기는 모델이 미리 정의해놓은 차원이 행, token의 갯수가 열로 이루어진다.</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/1eac480d-3019-430c-a5de-19b0a2cdbfe1" />
  </p>
</div>

<p><br /></p>

<p>이 embedding matrix는 모델이 어느 정도 단어까지를 문맥으로 볼 것인가에 대한 embedding_dimension x context_size의 크기로 변환되고</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/1a1cec7a-b9d4-4851-a0bd-bd677518d7f8" />
  </p>
</div>

<p><br /></p>

<p>변환된 embedding matrix의 마지막열 벡터를 사용하여 아래와 같은 unembedding matrix와 곱해(여기까지 logits) 확률(softmax 후)을 계산한다.
Unembedding matrix는 토큰갯수 x embedding_dimension의 크기를 가지고 있어 내적 후의 logits의 크기는 토큰갯수x1이다.</p>

<ul>
  <li>unembedding matrix: <strong>number of token</strong> x embedding dimension</li>
  <li>embedding matrix 마지막열 벡터: embedding dimension x <strong>1(last token)</strong></li>
  <li>둘의 내적: <strong>number of token</strong> x <strong>1(last token)</strong> -&gt; 토큰 갯수만큼의 logit값</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/4de1dbfd-5f9a-4cc7-964d-5f552b9c05c9" />
  </p>
</div>

<p><br /></p>

<p>softmax에서 temperature는 0에 가까울수록 가장 높은 logits값의 영향이 크도록(logit의 크기를 그대로 반영), 0과 멀어질수록 가장 높은 logits값의 영향이 작도록(다른 logits도 확률에 영향을 미치도록. 높은 logits값이 그대로 적용된다기 보다는 이 영향을 어느 정도 감쇠시켜서 다양한 토큰이 output으로 출력될 수 있음) 조절한다.</p>

<h1 id="attention">Attention</h1>
<p>Attention(여기서는 Self-Attention. Cross-Attention도 있는데 이건 나중에 다뤄보자)은 Query, Key, Value 를 통해 모델이 각 token들의 관계와 변화량을 계산하는 과정이다. 즉, 어떤 단어가 지금 나를 Attend(주목)하고 있는지 파악하는 과정.</p>

<ul>
  <li>Query는 나(현재 token)와 관계가 있는 token이 무엇인지 질문을 하는 역할</li>
  <li>Key는 어떤 token이 현재 token과 연관이 있는지 Query에 대한 답을 하는 역할</li>
  <li>Value는 Key에 담긴 의미를 나타내어, 기존 token의 embedding에 어떤 의미가 얼만큼 추가되어야 하는지 업데이트 하는 역할</li>
</ul>

<p>Q와 K는 관계(연관성)를 계산하기 위한 것이고, V는 실제 정보를 전달하기 위한 것이다.</p>

<p>$W_Q$ 는 embedding space에 있는 token vector를 query space로 옮겨주는 역할을 한다. $W_K$, $W_V$  는 key space, value space로 옮겨주는 역할을 한다.</p>

\[E_{1} \times W_Q = Q_{1}, \quad E_{1} \times W_K = K_{1}, \quad E_{1} \times W_V = V_{1}\]

<p>$Q_{1}$ 은 전체 토큰에 대해 만들어진 Query 행렬이며 각 행인 $i$번째 토큰의 Query 벡터는 $Q_{1}[i]$ 으로 표현할 수 있을 것(예시일뿐)이다. 즉, $Q_{1}$ 자체는 하나의 토큰을 나타내는 것이 아니라 모든 토큰에 대한 Query 벡터들의 집합이다. K, V도 마찬가지.</p>

<p><strong>예시</strong></p>
<ul>
  <li>Q1: 각 토큰에 대한 모든 query를 담고 있는 행렬</li>
  <li>Q[i]: “나는 이런 특징과 관련 있는 token의 정보를 받아오고 싶어.”</li>
  <li>K[j]: “나는 이런 특징을 가진 token이야.”</li>
  <li>Q[i]·K[j]: “그럼 우리 둘이 얼마나 관련 있음?”</li>
  <li>Value(V[j]): “내가 가진 실제 정보는 이거야. 네가 준 가중치(Q[i]·K[j])만큼 반영해줄게!”</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/4bbac9c7-99e5-42d0-b5f1-2dcad41b33ce" />
  </p>
</div>

<p><br /></p>

<p>token 사이의 관계는 Query와 Key의 내적으로 계산할 수 있고, 내적이 크면 두 token은 깊은 관계를 가지고 있는 것. 그리고 이 내적값은 softmax를 통해 확률값으로 변환할 수 있다.</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/421dc06f-6f37-4e21-b35f-82057dafe0dc" />
  </p>
</div>

<p><br /></p>

<p>추가적으로, Transformer는 RNN과 달리 전체 sequence를 한 번에 병렬로 처리할 수 있다는 것이 큰 장점이다. 하지만 decoder에서 학습 시, 뒤에 나올 단어를 이미 알면 학습에 의미가 없으므로 future token을 가린 상태(masking)로 Attention을 수행한다. 이때 Q와 K의 내적값에서 뒤의 단어에 해당하는 부분에 -Inf를 넣어 softmax 확률이 0이 되도록 하여, 해당 단어가 사용되지 않도록 한다.</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/d5fea525-8c8f-4063-9c76-81a83bc608ba" />
  </p>
</div>

<p><br /></p>

<p>Value vector는 문맥 정보를 잘 담은 최종 token을 만들기 위해 다른 token에서 더해져야 할 정보를 담고 있는 벡터로, 기존 embedding이 얼마나 변해야 하는지를 나타낸다. Q,K가 내적된 후 softmax를 통해 확률값으로 변하면 이 확률값은 Value를 얼만큼 적용할지를 결정하는 가중치 역할을 하게 되고,
이 Value가 기존 embedding vector를 얼만큼 바꿀지를 결정하게 되는것이다.</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/ce1e22d4-c098-4cc2-940d-99845af0201b" />
  </p>
</div>

<p><br /></p>

<p>따라서, 이와 같이 모델에 의해 학습되는 $W_{Q1}$, $W_{K1}$, $W_{V1}$ -&gt; $Q_{1}$, $K_{1}$, $V_{1}$ 을 한 세트로 한 가지 관점에서 token의 관계를 파악하는 것을 Single Head Attention이라고 하고, 각기 다른 $W_{Q}$, $W_{K}$ , $W_{V}$ 를 병렬로 배치하여
다양한 문맥과 관계를 학습할 수 있도록 하는 것을 Multi-Head Attention이라고 한다.</p>

<p>Sequence of token embeddings: $E = [e_1, e_2, \dots, e_n]$</p>

<p>Dimension: $d_{\text{model}}$</p>

<table>
  <thead>
    <tr>
      <th style="text-align: center">Single</th>
      <th style="text-align: center">Multi</th>
    </tr>
  </thead>
  <tbody>
    <tr>
      <td style="text-align: center"><img src="https://github.com/user-attachments/assets/cb9a42da-afe5-4d98-b526-a5452af6df17" width="70%" /></td>
      <td style="text-align: center"><img src="https://github.com/user-attachments/assets/fcac310c-d051-4e7e-91aa-bae81bf329e3" width="70%" /></td>
    </tr>
  </tbody>
</table>

<h1 id="multilayer-perceptron">MultiLayer Perceptron</h1>
<p>Attention이 문맥 정보를 계산한다면, MLP는 문맥 정보를 기반으로 중요한 feature를 더 강화한다. MLP는 Attention의 output을 입력으로 받아 linear, R(G)eLU, linear 연산 후 입력을 다시 더해주는(skip connection) 방법을 사용한다.</p>

<p>예를 들어, Attention의 output에서 마이클 조던이라는 이름에 대한 정보를 담고 있는 임베딩 벡터하나를 MLP에 넣으면 여러 연산을 통해 농구와 관련된 임베딩 벡터가 더 강화되고 이걸 다시 마이클 조던이라는 input과 더해서 마이클 조던이 농구와 관련되었다는 정보를 더 강화하게 된다. 마이클 조던과 농구가 관련 있다는 것은 이미 Attention에서 문맥 정보를 파악하며 이루어졌고, MLP는 이렇게 중요한 정보를 비선형 연산을 통해 더 강화하는 역할을 한다. 이런 과정이 모든 토큰 벡터에 똑같이 적용된다.</p>

<p><br /></p>
<div align="center">
  <p>
    <img style="width: 70%;" alt="image" src="https://github.com/user-attachments/assets/66ebf2d2-4c82-46c0-a843-a405189dffdc" />
  </p>
</div>

<p><br /></p>

<p>LLM은 이렇게 text를 입력으로 받아 Attention으로 문맥을 파악하고, MultiLayer Perceptrion에서 중요한 토큰은 더 강화하고, 중요하지 않은 토큰은 약화시키는 과정이 여러번 반복되어 다음 단어를 예측하는 구조로 동작한다.</p>

<p>특히, 요즘에는 MLP의 각 차원이 하나의 의미만 담당하는 것이 아니라(마이클 조던이 마이클 조던의 이름을 넘어 농구, 시카고 불스, 나이키, 23번, 은퇴 등 훨씬 더 많은 정보를 담고 있도록), 여러 의미가 희소하게 겹쳐져 인코딩되는 ‘Super Position’이라는 개념이 주목을 받아, LLM이 제한된 파라미터로 방대한 지식을 저장할 수 있도록 하고, 이 구조를 해석하려는 interpretability 연구가 활발히 이루어지고 있다.</p>

<p>Transformer와 MLP 학습에는 아래 자료를 참고했다.</p>

<ul>
  <li><a href="https://www.youtube.com/watch?v=g38aoGttLhI">트랜스포머, ChatGPT가 트랜스포머로 만들어졌죠. - DL5</a></li>
  <li><a href="https://www.youtube.com/watch?v=_Z3rXeJahMs">그 이름도 유명한 어텐션, 이 영상만 보면 이해 완료! - DL6</a></li>
  <li><a href="https://www.youtube.com/watch?v=zHQLPJ8-9Qc">수많은 정보는 LLM 모델 속 어디에 저장되어있는걸까? - DL 7</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Interpretability" /><summary type="html"><![CDATA[머신러닝과 비교하여 딥러닝 모델의 가장 큰 문제는 모델이 왜 이런 결과를 만들어냈는지 해석할 수 없다는 것이고 이를 ‘Black Box Problem’이라고 한다. 그리고 이 Black Box 문제를 해결해보고자 예전부터 XAI(eXplainable AI)에 관한 연구가 활발한데 XAI에서 사용하는 방법은 2가지로 나눌 수 있다. 해석가능성(Interpretability): 엔지니어의 관점에서 모델의 구조를 기반으로 알고리즘이 어떻게 변화하여 해당 결과를 도출해내는지 해석할 수 있는 능력 설명가능성(Explainability): 해석가능성에서 더 나아가 사회과학, HCI 등의 요소를 결합하여 일반적인 사용자들도 이해할 수 있도록 설명하는 능력]]></summary></entry><entry><title type="html">EfficientNet, EfficientDet 정리</title><link href="https://kimyeonz.github.io/blog/detection/EfficientNet_and_EfficientDet/" rel="alternate" type="text/html" title="EfficientNet, EfficientDet 정리" /><published>2023-11-20T00:00:00+00:00</published><updated>2023-11-20T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/EfficientNet_and_EfficientDet</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/EfficientNet_and_EfficientDet/"><![CDATA[<p>Computer vision 분야의 발달로 네트워크의 층을 더해 정확도를 높이면서 연산량을 줄여 속도는 높이는 방법들이 고안되었다. 이 과정속에서 네트워크의 깊이와 높이, 이미지의 크기, 필터의 채널 수 등 다양한 파라미터들이
실험적으로 사용되었는데 EfficientNet과 EfficientDet에서는 compoung scaling을 사용해 해당 네트워크에 최적의 파라미터들을 찾아 성능을 개선시켰다. 이번 포스팅에서는 EfficientNet과 EfficientDet에 대해 간단히
정리해보자.</p>

<h1 id="efficientnet">EfficientNet</h1>
<p>EfficientNet에서는 기본 backbone구조에서 네트워크의 넓이(필터 수), 깊이, 이미지 resolution을 독립적으로 scaling하거나 compound scaling을 수행하여 여러 네트워크에 대한 성능을 비교했다. 여기서 compound scaling이란,
각각의 파라미터를 scaling하고 이를 조합하여 최적의 파라미터 조합을 찾는 것으로 네트워크의 깊이, 필터 수, 이미지 resolution 크기가 파라미터로 사용된다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6d9564f5-2c1d-4985-812d-3a2cc7fbf109" />
  </p>
</div>

<p><br /></p>

<p>Figure 3.을 참고하면 각각의 파라미터들을 독립적으로 scaling을 수행할 때 resolution을 제외하고는 연산량이 증가해도 결국 정확도의 증가폭이 작아지는 구간이 생긴다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/fbc8b1b5-55a1-450e-97dc-75fdcbb287d0" />
  </p>
</div>

<p><br /></p>

<p>이를 해결하기 위해 EfficientNet에서는 모든 파라미터를 scaling하는 compound scaling 기법을 사용했다. Depth, width, resolution 이 3가지 factor를 동시에 고려하는데 아래의 식처럼 최초에는
$\sigma$를 1로 고정하고 grid search를 기반으로 최적의 $\alpha, \beta, \gamma$ 값을 찾는다(EfficientNetB0는 1.2, 1.1, 1.15).</p>

<p>다음으로 $\alpha, \beta, \gamma$를 고정하고 $\sigma$ 값을 증가시켜가면서 EfficientB1 ~ B7까지 scale 조합을 구성한다. 특히, width와 resolution을 경우 factor가 2배가 되면 FLOPS가 4배가 되기대문에
너무 커지지않도록 제곱을 통해 scale을 조정했다. 이러한 방법으로 depth, width, resolution에 따른 FLOPS 변화를 기반으로 최적의 식을 도출한다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/94200458-8a65-469a-9354-1d17a4d139e4" />
  </p>
</div>

<p><br /></p>

<p>이런 다양한 실험을 기반으로 EfficientNet은 ImageNet Classification에서 기존 SOTA 모델들에 비해 정확도와 속도 모두 향상된 performance를 기록했다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/8008b4da-84be-412c-8baf-4cfea84fa324" />
  </p>
</div>

<p><br /></p>

<h1 id="efficientdet">EfficientDet</h1>
<h2 id="bi-fpn">Bi-FPN</h2>
<p>RetinaNet이 FPN을 성공적으로 사용함으로써 FPN에 대한 연구가 매우 활발해졌다. EfficientDet은 EfficientNet에 Bi-FPN을 접목시켜 Detection task를 수행하는 모델이다.</p>

<p>FPN이 top-down 과정에서 bottom-up의 피처맵을 결합하여 예측했다면, Bi-FPN은 top-down뿐만 아니라 다시 bottom-up을 수행해 이 과정에서 이전 bottom-up의 피처맵과 top-down에서 결합된 피처맵을 모두 사용한다.
FPN의 아이디어는 높은 추상화 정보를 가지고 있는 피처맵에 낮은 추상화 정보를 가지고 있는 피처맵을 결합하여 공간 정보를 유지시키는 것이었다. Bi-FPN은 여기서 한발 더 나아가 낮은 추상화 정보의 피처맵도 높은 추상화 정보의 피처맵이
있다면 예측을 더 잘 수행할 것이라는 아이디어다. 즉, FPN보다 공간 정보를 훨씬 더 많이 사용하여 예측을 수행한다는 것이다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/104659c0-5305-47a2-9a34-ea843a1914d2" />
  </p>
</div>

<p><br /></p>

<p>이 Bi-FPN block이 하나만 있는 것이 아니라 여러개가 존재하고, 각각의 피처맵들을 더할 때 가중치가 너무 커지지 않도록 Weighted Feature Fusion을 사용한다. Weighted Feature Fusion의 기법으로는
softmax, unbounded fusion을 사용해봤지만 Fast normalized fusion을 사용했을 때 연산량의 측면에서 가장 효율적이었다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="200" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/a8b35e94-1bda-407b-ada0-2e30cd784f02" />
    <br />
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/ebdfcd75-de5b-48e7-a46e-2e614b91ba72" />
  </p>
</div>
<p><br /></p>

<h2 id="compound-scaling">Compound Scaling</h2>
<p>EfficientDet은 Backbone, Neck, Head에 모두 compound scaling을 적용했는데 backbone은 EfficientNet scaling을 그대로 적용했고, Bi-FPN network에는 $D_{bifpn} = 3 + \sigma$로
기본 반복 block을 3개로 적용하여 scaling했다.</p>

<p>Width(채널 수)는 $W_{bifpn} = 64 * (1.35 \sigma)$, Prediction network의 depth는 $D_{box} = D_{class} = 3 + [\sigma / 3]$, 이미지 resolution은 $R_{input} = 512 + \sigma * 128$
로 scaling했다.</p>

<p>여러 실험을 통해 최종적으로는 네트워크의 아키텍처에 따라 아래와 같이 compound scaling 되었다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/7e05ed67-7e79-44c3-9a90-f63b33abd945" />
  </p>
</div>

<p><br /></p>

<p>EfficientDet은 EfficientNet을 기반으로 compound scaling을 사용했고, FPN을 개선한 Bi-FPN, activation으로 SiLU, loss로 Focal Loss, NMS를 조금더 보완한 Soft-NMS(기존 NMS가 근처의 detection된 박스까지 제거하는 것을
보완하여 많이 겹치는 박스에 가중치를 크게 부여해 confidence score를 낮춰서 겹치는 다른 박스도 어느 정도 유지시킬 수 있도록 하는 기법)기법을 사용하여 성능을 향상시켰다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/8674f097-ef44-4fb0-b53b-7cea7cbad70a" />
  </p>
</div>

<p><br /></p>

<p>하지만 여전히, small object 탐지는 어렵다는 한계가 존재하긴한다.</p>

<h1 id="reference">Reference</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1905.11946.pdf">EfficientNet Paper</a></li>
  <li><a href="https://arxiv.org/pdf/1911.09070.pdf">EfficientDet Paper</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[Computer vision 분야의 발달로 네트워크의 층을 더해 정확도를 높이면서 연산량을 줄여 속도는 높이는 방법들이 고안되었다. 이 과정속에서 네트워크의 깊이와 높이, 이미지의 크기, 필터의 채널 수 등 다양한 파라미터들이 실험적으로 사용되었는데 EfficientNet과 EfficientDet에서는 compoung scaling을 사용해 해당 네트워크에 최적의 파라미터들을 찾아 성능을 개선시켰다. 이번 포스팅에서는 EfficientNet과 EfficientDet에 대해 간단히 정리해보자.]]></summary></entry><entry><title type="html">Fast &amp;amp; Faster R-CNN 정리</title><link href="https://kimyeonz.github.io/blog/detection/Fast_and_Faster_RCNN/" rel="alternate" type="text/html" title="Fast &amp;amp; Faster R-CNN 정리" /><published>2023-11-20T00:00:00+00:00</published><updated>2023-11-20T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/Fast_and_Faster_RCNN</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/Fast_and_Faster_RCNN/"><![CDATA[<p>R-CNN이 다른 모델들에 비해 높은 mAP를 기록했고 이후 Fast-RCNN, Faster R-CNN, Mask R-CNN 등 여러 모델들이 R-CNN을 develop하여 object detection 분야에서 더 많은 발전을 이루어냈다. 이번 포스팅에서는 R-CNN 이후의 모델인 Fast R-CNN과 Faster R-CNN에 대해 간단하게 정리해보자.</p>

<h1 id="spatial-pyramid-pooling-netsppnet">Spatial Pyramid Pooling Net(SPPNet)</h1>
<p>Fast &amp; Faster R-CNN을 더 잘 이해하기위해 먼저, SPPNet에 대해 알아보자.</p>

<p>R-CNN은 selective search와 CNN, SVM을 통해 object detection을 수행했는데 region proposals, feature extraction, detection이 모두 다른 네트워크에서 수행되어 학습 및 추론 시간이 느리다는 단점을 가지고있다.
또한, 2000개의 region proposals이 CNN에 입력되기 위해서는 모든 region이 고정된 크기의 벡터여야하는데 이를 위해 crop이나 warp를 사용해 사이즈를 조절했다. Crop/warp를 사용할 경우 실제 region과는 조금은 다른 이미지가 만들어지기때문에 성능에도 영향이 있다.</p>

<p>SPPNet은 이러한 문제점을 개선해서 2000개의 region proposals 이미지를 전부 CNN에 통과시키는 것이 아니라 CNN은 원본 이미지만 통과시켜 피처맵을 만들고 selective search로 나온 region을 이 피처맵과 맵핑시켜 학습하는 방법을 사용했다. 하지만, 피처맵과 맵핑된 region을 고정된 크기의 벡터로 flatten시켜야 분류기에 넣어 detection을 수행할 수 있는데 모든 region의 크기가 제각각이기때문에 이를 고정된 크기의 벡터로 만드는 것이 어려웠다.</p>

<p>SPPNet에서는 이를 해결하기위해 SPP Layer를 만들어서 Flatten을 시키기전에 어떤 크기의 region이 들어와도 분류기에 입력하기 전에 고정된 크기의 벡터로 변환했다. 아래 그림처럼 다양한 크기의 피처맵이 들어오면 해당 피처맵을 여러개의 분면으로 쪼개고 이를 합쳐서 고정된 크기의 벡터로 만들어 분류기에 입력으로 사용하는 것이다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/52ba1d69-7923-400d-bcf8-624bd5cca7ba" />
  </p>
</div>

<p><br /></p>

<p>이를 통해 성능의 개선 뿐만 아니라 detection 수행 시간도 크게 단축됐다.</p>

<h1 id="fast-r-cnn">Fast R-CNN</h1>
<p>Fast R-CNN은 SPPNet을 조금 더 보완시켜 SPP layer를 RoI pooling layer로 바꾼 네트워크라고 할 수 있다. 또한, 위에서 언급한 R-CNN의 문제를 여러가지 방법을 통해 해결했는데 먼저 SVM을 softmax 네트워크로 변환시켰고
classification과 regression을 혼합하여 사용하는 multi-task loss 함수를 통해 end-to-end network(RoI proposals 제외)를 구축했다.</p>

<p>SPPNet이 다양한 크기의 피처맵을 미리 정해놓은 여러개의 분면으로 쪼개 고정된 크기의 벡터로 만들었다면, Fast R-CNN은 RoI pooling을 통해 어떤 크기의 피처맵이 들어와도 모두 고정된 크기로 max pooling하는 기법을 사용한다.
보통 7 x 7 크기로 pooling 시키는데 만약 피처맵의 크기가 7의 배수가 아니라면 보간이나 이미지를 resizing하여 크기를 맞춘다.</p>

<p>손실함수로는 multi task loss를 사용하고 특히 regression loss에는 smooth $L_1$을 적용해서 R-CNN과 SPPNet에서 사용된 $L_2$ loss보다 outliers에 덜 민감하도록 하고, 
loss가 1보다 작으면 loss를 더 작게해서 큰 loss에 더 집중할 수 있도록 한다. $\lambda$는 classification과 box regression loss의 balance를 맞추는 용도로 사용된다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/a12312fb-2d42-4850-bf1a-078d715be168" />
  </p>
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/ecf36288-8f1f-4017-89a8-b251ef7459c7" />
  </p>
  <p>
  <img width="450" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6b7d4f67-7234-4a08-9a28-a5fb794dc05a" />
  </p>
</div>
<p><br /></p>

<p>Fast R-CNN의 학습 과정을 정리하면 다음과 같다.</p>

<p>1) 원본 이미지를 CNN에 통과시켜 feature extraction을 수행</p>

<p>2) selective search를 통해 나온 2000개의 region proposals을 원본 이미지의 피처맵과 맵핑</p>

<p>3) 맵핑된 다양한 크기의 region proposals 피처맵에 7x7 RoI pooling 적용</p>

<p>4) RoI pooling을 통해 나온 output으로 detection 수행(multi task loss)</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/19f94db4-a3da-4c4f-b07a-b4af21e67fa2" />
  <br />
  <img width="300" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/95ca23fc-703a-4540-9348-901ee76d56e8" />\
  </p>
</div>
<p><br /></p>

<p>결과적으로 Fast R-CNN은 속도나 정확도 측면 모두에서 기존 R-CNN보다 크게 개선된 mAP를 기록했다.</p>

<h1 id="faster-r-cnn">Faster R-CNN</h1>
<p>Fast R-CNN은 CNN, RoI pooling, 분류기를 결합해 특징 추출과 detection을 하나의 네트워크에서 수행하여 R-CNN보다 훨씬 빠른 속도로 detection이 가능했지만 여전히 region proposals에는 selective search 알고리즘을 사용하여 완벽한 end-to-end network를
구축하진 못했고 one stage model보다 속도 측면에서 좋지 않은 performance를 보였다. Faster R-CNN은 이러한 문제를 해결하기위해 Fast R-CNN에 RPN(Region Proposals Network)을 결합하여 selective search를 대체한
완벽한 end-to-end network를 구축했다.</p>

<p>End-to-End Network가 구축되면 classification, box regression 뿐만 아니라 region에 대한 back propagation도 가능해지기때문에 2000개로 한정된 selective search 알고리즘보다 더 효율적인 proposals이 가능하다.
그렇다면, 이미지에서 어떻게 selective search와 같이 region을 예상하여 제안할 수 있을까?</p>

<p>여기에서 사용된 개념이 anchor box이다. Anchor box는 한 픽셀당 고정된 여러 스케일의 boxes를 사용하여 region proposals을 수행한다. 아래 그림처럼 각각의 피처맵의 한 픽셀당 여러 ratio와 크기를 가진 k개의 anchor box가 있다고 할 때 anchor box마다 object의 존재여부 scores와($2k$ scores) 4개의 box 좌표($4k$ scores)가 output으로 도출된다. Anchor boxes가 ground truth box와 일치한다의 기준은 IoU가 가장 높은 anchor나 0.7이상인 anchor를 positive, 0.3이하를 negative로 분류하고 그 외의 애매한 box는 학습에서 제외시킨다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/a37bfcfd-7fb1-44a9-83da-afe1129efbab" />
  </p>
</div>
<p><br /></p>

<p>손실 함수는 Fast R-CNN처럼 multi task loss를 사용하지만 anchor box와 관련된 계수들이 추가되었다. $p_i$는 anchor box내 객체가 object일 확률이고, $p_{i}^{\star}$는 해당 object가 ground trurh와 일치하면 positive로 1, 일치하지 않으면 negative 0으로 취급한다. Box regression은 anchor box가 클래스를 올바르게 예측한 즉, positive에 대해서만 수행하고 $t_i$는 anchor box와 모델이 예측한 bbox와의 차이, $t_{i}^{\star}$는 anchor box와 ground truth간의 차이이다. Faster R-CNN의 손실함수에서 특이한 점은 예측 bounding box와 ground truth box의 차이를 anchor box와 각각의 prediction box와의 차이를 사용하여 계산한다는 것이다.</p>

<p>이 방법은 anchor box를 참고해서 anchor를 기준으로 GT와 predicted box의 차이가 비슷할수록 두 박스의 거리가 가까울 것이라는 아이디어이다. 객체의 존재여부를 하나도 모르는 bbox를 생성하고 조정하는 것보다 positive anchor를 참고하여 bbox를 GT에 가깝게 조금씩 조정하는 것이 더 효율적인 방법이 된다. $N_{cls}$는 positive와 negative anchor의 비율을 동일하게 가져가기 위한 정규화 파라미터고 $N_{reg}$는 박스 갯수를 정규화한 값이다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="450" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/9911a43e-0cde-4a93-8700-cd1a62dacf80" />
  <br />
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6db411c0-3104-4626-bcbc-d353de5b57e6" />
  </p>
</div>
<p><br /></p>

<p>Faster R-CNN의 학습 과정을 정리하면 다음과 같다.</p>

<p>1) 원본 이미지를 CNN을 통과시켜 특징 추출</p>

<p>2) 각각의 피처맵에 대해 픽셀당 여러 스케일을 가진 anchor box를 그리고 각각의 boxes를 RoI pooling으로 고정된 크기의 벡터로 변환</p>

<p>3) Anchor boxes에 대해 object 존재여부와 ground truth와의 일치여부를 계산하고, 일치한다면 anchor box, predicted box, ground truth box를 통해 box regression 수행</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/129d93c3-c0d4-4c8b-a7f2-e78b70484aef" />
  <br />
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/bcc38a5d-40ec-4d16-9b0a-138a76217953" />
  </p>
</div>
<p><br /></p>

<p>Faster R-CNN은 selective search 보다 anchor boxes를 활용한 RPN을 사용했을 때 성능이 크게 향상됐고, Fast R-CNN보다 개선된 모델임을 알 수 있다.</p>

<h1 id="구현">구현</h1>
<p>Pytorch로 Faster R-CNN model을 사용해보자. 데이터는 Roboflow에서 제공하는 <a href="https://universe.roboflow.com/yinguo/soccer-data">soccer dataset</a>을 사용했다. Class는 공을 소유한 사람과 소유하지 않은 사람으로 구분된다.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># engine 라이브러리 사용을 위한 git clone
</span><span class="err">!</span><span class="n">git</span> <span class="n">clone</span> <span class="n">https</span><span class="p">:</span><span class="o">//</span><span class="n">github</span><span class="p">.</span><span class="n">com</span><span class="o">/</span><span class="n">pytorch</span><span class="o">/</span><span class="n">vision</span><span class="p">.</span><span class="n">git</span>
<span class="o">%</span><span class="n">cd</span> <span class="n">vision</span>
<span class="err">!</span><span class="n">git</span> <span class="n">checkout</span> <span class="n">v0</span><span class="p">.</span><span class="mf">3.0</span>

<span class="err">!</span><span class="n">cp</span> <span class="n">references</span><span class="o">/</span><span class="n">detection</span><span class="o">/</span><span class="n">utils</span><span class="p">.</span><span class="n">py</span> <span class="p">..</span><span class="o">/</span>
<span class="err">!</span><span class="n">cp</span> <span class="n">references</span><span class="o">/</span><span class="n">detection</span><span class="o">/</span><span class="n">transforms</span><span class="p">.</span><span class="n">py</span> <span class="p">..</span><span class="o">/</span>
<span class="err">!</span><span class="n">cp</span> <span class="n">references</span><span class="o">/</span><span class="n">detection</span><span class="o">/</span><span class="n">coco_eval</span><span class="p">.</span><span class="n">py</span> <span class="p">..</span><span class="o">/</span>
<span class="err">!</span><span class="n">cp</span> <span class="n">references</span><span class="o">/</span><span class="n">detection</span><span class="o">/</span><span class="n">engine</span><span class="p">.</span><span class="n">py</span> <span class="p">..</span><span class="o">/</span>
<span class="err">!</span><span class="n">cp</span> <span class="n">references</span><span class="o">/</span><span class="n">detection</span><span class="o">/</span><span class="n">coco_utils</span><span class="p">.</span><span class="n">py</span> <span class="p">..</span><span class="o">/</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="err">!</span><span class="n">pip</span> <span class="n">install</span> <span class="o">-</span><span class="n">q</span> <span class="n">torch</span><span class="o">==</span><span class="mf">1.13</span><span class="p">.</span><span class="mi">0</span> <span class="n">torchvision</span><span class="o">==</span><span class="mf">0.14</span><span class="p">.</span><span class="mi">0</span> <span class="c1"># 런타임 다시 시작
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">json</span>
<span class="kn">import</span> <span class="nn">cv2</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">os</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
<span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">from</span> <span class="nn">pycocotools.coco</span> <span class="kn">import</span> <span class="n">COCO</span>
<span class="kn">from</span> <span class="nn">PIL</span> <span class="kn">import</span> <span class="n">Image</span>
<span class="kn">import</span> <span class="nn">time</span>
<span class="kn">import</span> <span class="nn">transforms</span> <span class="k">as</span> <span class="n">T</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/'</span>
<span class="n">TR_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/train/'</span>
<span class="n">VAL_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/valid/'</span>
<span class="n">TEST_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/test/'</span>
<span class="n">TR_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/train/_annotations.json'</span>
<span class="n">VAL_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/valid/_annotations.json'</span>
<span class="n">TEST_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/test/_annotations.json'</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Custom Datset 정의
</span><span class="k">class</span> <span class="nc">SoccerDataset</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">utils</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="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_path</span><span class="p">,</span> <span class="n">label_path</span><span class="p">,</span> <span class="n">transforms</span><span class="p">):</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">data_path</span> <span class="o">=</span> <span class="n">data_path</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">label_path</span> <span class="o">=</span> <span class="n">label_path</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span> <span class="o">=</span> <span class="n">transforms</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">imgs</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">data_path</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.jpg'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">labs</span> <span class="o">=</span> <span class="n">COCO</span><span class="p">(</span><span class="n">label_path</span><span class="p">)</span> <span class="c1"># COCO를 사용하여 라벨을 쉽게 불러올 수 있다
</span>
        <span class="c1"># 이미지 전체 id
</span>        <span class="n">all_img_ids</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">.</span><span class="n">getImgIds</span><span class="p">()</span> <span class="c1"># 이미지 id 전체 가져오기
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">img_ids</span> <span class="o">=</span> <span class="p">[]</span>

        <span class="k">for</span> <span class="n">idx</span> <span class="ow">in</span> <span class="n">all_img_ids</span><span class="p">:</span>
            <span class="n">annotations_ids</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">.</span><span class="n">getAnnIds</span><span class="p">(</span><span class="n">imgIds</span><span class="o">=</span><span class="n">idx</span><span class="p">,</span> <span class="n">iscrowd</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span> <span class="c1"># 해당 img_id에 일치하는 annotation들
</span>            <span class="k">if</span> <span class="nb">len</span><span class="p">(</span><span class="n">annotations_ids</span><span class="p">)</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># 만약 list가 0이면 해당 image에는 annotation이 없는 것
</span>                <span class="k">print</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span>
            <span class="k">else</span><span class="p">:</span>
                <span class="bp">self</span><span class="p">.</span><span class="n">img_ids</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">__getitem__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">idx</span><span class="p">):</span>
        <span class="n">image</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">load_image</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span>
        <span class="n">lab</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">load_annotatinos</span><span class="p">(</span><span class="n">idx</span><span class="p">)</span>
        
        <span class="n">boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">lab</span><span class="p">[:,:</span><span class="mi">4</span><span class="p">])</span>
        <span class="n">labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">lab</span><span class="p">[:,</span><span class="mi">4</span><span class="p">],</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>
        <span class="n">area</span> <span class="o">=</span> <span class="p">(</span><span class="n">boxes</span><span class="p">[:,</span> <span class="mi">3</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">])</span> <span class="o">*</span> <span class="p">(</span><span class="n">boxes</span><span class="p">[:,</span> <span class="mi">2</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">])</span> <span class="c1"># width * height
</span>        <span class="n">iscrowd</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="n">boxes</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">],),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span> <span class="c1"># 해당 라벨들은 군집화되어있지 않음. 따라서 box수만큼 0표시
</span>
        <span class="n">target</span> <span class="o">=</span> <span class="p">{}</span> <span class="c1"># engine 라이브러리를 사용하기위해선 항상 아래 형식의 target을 만들어야한다.
</span>        <span class="n">target</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">]</span> <span class="o">=</span> <span class="n">boxes</span>
        <span class="n">target</span><span class="p">[</span><span class="s">'labels'</span><span class="p">]</span> <span class="o">=</span> <span class="n">labels</span>
        <span class="n">target</span><span class="p">[</span><span class="s">"image_id"</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">img_ids</span><span class="p">[</span><span class="n">idx</span><span class="p">])</span>
        <span class="n">target</span><span class="p">[</span><span class="s">'area'</span><span class="p">]</span> <span class="o">=</span> <span class="n">area</span>
        <span class="n">target</span><span class="p">[</span><span class="s">"iscrowd"</span><span class="p">]</span> <span class="o">=</span> <span class="n">iscrowd</span>

        <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="c1"># 이렇게 넘어온 key, value 값으로 다시 리턴해줌
</span>            <span class="c1"># bboxes, labels라는 key로 transformed dict 리턴
</span>            <span class="n">transformed</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span><span class="p">(</span><span class="n">image</span> <span class="o">=</span> <span class="n">image</span><span class="p">,</span> <span class="n">bboxes</span> <span class="o">=</span> <span class="n">boxes</span><span class="p">,</span> <span class="n">labels</span><span class="o">=</span><span class="n">labels</span><span class="p">)</span>
            <span class="n">image</span> <span class="o">=</span> <span class="n">transformed</span><span class="p">[</span><span class="s">'image'</span><span class="p">]</span>
            <span class="n">target</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">transformed</span><span class="p">[</span><span class="s">'bboxes'</span><span class="p">])</span> <span class="c1"># 변환된 box로 정의
</span>            <span class="n">target</span><span class="p">[</span><span class="s">'labels'</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">transformed</span><span class="p">[</span><span class="s">'labels'</span><span class="p">])</span>
            
        <span class="k">return</span> <span class="n">T</span><span class="p">.</span><span class="n">ToTensor</span><span class="p">()(</span><span class="n">image</span><span class="p">,</span> <span class="n">target</span><span class="p">)</span> <span class="c1"># ToTensor는 이미지를 정규화하고 (C,H,W) 형식으로 만듬
</span>    
    <span class="k">def</span> <span class="nf">load_image</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">img_idx</span><span class="p">):</span>
        <span class="n">image_info</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">.</span><span class="n">loadImgs</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">img_ids</span><span class="p">[</span><span class="n">img_idx</span><span class="p">])[</span><span class="mi">0</span><span class="p">]</span> <span class="c1"># 일치하는 이미지 정보를 가져옴
</span>        <span class="n">img</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">imread</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">data_path</span> <span class="o">+</span> <span class="n">image_info</span><span class="p">[</span><span class="s">'file_name'</span><span class="p">])</span>
        <span class="n">img</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">cvtColor</span><span class="p">(</span><span class="n">img</span><span class="p">,</span> <span class="n">cv2</span><span class="p">.</span><span class="n">COLOR_BGR2RGB</span><span class="p">)</span>

        <span class="k">return</span> <span class="n">img</span>

    <span class="k">def</span> <span class="nf">load_annotatinos</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">img_idx</span><span class="p">):</span>
        <span class="n">annot_ids</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">.</span><span class="n">getAnnIds</span><span class="p">(</span><span class="n">imgIds</span><span class="o">=</span><span class="bp">self</span><span class="p">.</span><span class="n">img_ids</span><span class="p">[</span><span class="n">img_idx</span><span class="p">],</span> <span class="n">iscrowd</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span> <span class="c1"># 해당 img id와 일치하는 annot_id
</span>        <span class="n">annots</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="mi">0</span><span class="p">,</span><span class="mi">5</span><span class="p">))</span> <span class="c1"># box좌표와 category_id를 넣을 array
</span>
        <span class="k">if</span> <span class="nb">len</span><span class="p">(</span><span class="n">annot_ids</span><span class="p">)</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
            <span class="k">print</span><span class="p">(</span><span class="s">'No annotations in this image'</span><span class="p">)</span>

        <span class="n">coco_annots</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">.</span><span class="n">loadAnns</span><span class="p">(</span><span class="n">annot_ids</span><span class="p">)</span> <span class="c1"># 라벨이 존재하면 모든 라벨 정보 저장
</span>        <span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">coco_annots</span><span class="p">):</span>
            <span class="n">annot</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span><span class="mi">5</span><span class="p">))</span>
            <span class="n">annot</span><span class="p">[</span><span class="mi">0</span><span class="p">,:</span><span class="mi">4</span><span class="p">]</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="s">'bbox'</span><span class="p">]</span>
            <span class="n">annot</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span><span class="mi">4</span><span class="p">]</span> <span class="o">=</span>  <span class="n">x</span><span class="p">[</span><span class="s">'category_id'</span><span class="p">]</span> <span class="c1"># 학습시 loss_box_reg가 0인거면 라벨링이 잘못된것
</span>            <span class="n">annots</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">annots</span><span class="p">,</span> <span class="n">annot</span><span class="p">,</span> <span class="n">axis</span> <span class="o">=</span> <span class="mi">0</span><span class="p">)</span>
        
        <span class="c1"># w,h를 x2,y2형식으로 변환
</span>        <span class="n">annots</span><span class="p">[:,</span><span class="mi">2</span><span class="p">]</span> <span class="o">=</span> <span class="n">annots</span><span class="p">[:,</span><span class="mi">0</span><span class="p">]</span> <span class="o">+</span> <span class="n">annots</span><span class="p">[:,</span><span class="mi">2</span><span class="p">]</span>
        <span class="n">annots</span><span class="p">[:,</span><span class="mi">3</span><span class="p">]</span> <span class="o">=</span> <span class="n">annots</span><span class="p">[:,</span><span class="mi">1</span><span class="p">]</span> <span class="o">+</span> <span class="n">annots</span><span class="p">[:,</span><span class="mi">3</span><span class="p">]</span>
        <span class="c1">#print(annots)
</span>        <span class="k">return</span> <span class="n">annots</span>

    <span class="k">def</span> <span class="nf">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">imgs</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">albumentations</span> <span class="k">as</span> <span class="n">A</span>

<span class="k">def</span> <span class="nf">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="p">):</span>
    <span class="n">transforms</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">if</span> <span class="n">train</span><span class="p">:</span>
        <span class="n">transforms</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">A</span><span class="p">.</span><span class="n">HorizontalFlip</span><span class="p">(</span><span class="mf">0.5</span><span class="p">))</span>
        <span class="n">transforms</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">A</span><span class="p">.</span><span class="n">VerticalFlip</span><span class="p">(</span><span class="mf">0.5</span><span class="p">))</span>
    <span class="k">return</span> <span class="n">A</span><span class="p">.</span><span class="n">Compose</span><span class="p">(</span><span class="n">transforms</span><span class="p">,</span> <span class="n">bbox_params</span><span class="o">=</span><span class="n">A</span><span class="p">.</span><span class="n">BboxParams</span><span class="p">(</span><span class="nb">format</span><span class="o">=</span><span class="s">'pascal_voc'</span><span class="p">,</span> <span class="n">label_fields</span><span class="o">=</span><span class="p">[</span><span class="s">'labels'</span><span class="p">]))</span>
    <span class="c1"># label_fields는 호출할 때 입력한 key와 맞아야함. key이름을 labels로 했으니까 label_field로 labels
</span>    <span class="c1"># 이미 x2,y2형식으로 바꿨으니까 coco가 아닌 pascal 형식으로 리턴
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 모델 정의
</span><span class="kn">import</span> <span class="nn">torchvision</span>
<span class="kn">from</span> <span class="nn">torchvision.models.detection.faster_rcnn</span> <span class="kn">import</span> <span class="n">FastRCNNPredictor</span><span class="p">,</span> <span class="n">AnchorGenerator</span>

<span class="n">model</span> <span class="o">=</span> <span class="n">torchvision</span><span class="p">.</span><span class="n">models</span><span class="p">.</span><span class="n">detection</span><span class="p">.</span><span class="n">fasterrcnn_resnet50_fpn</span><span class="p">(</span><span class="n">weights</span> <span class="o">=</span> <span class="s">'DEFAULT'</span><span class="p">)</span> <span class="c1"># 사전학습된 가중치 그대로 사용
</span><span class="n">num_classes</span> <span class="o">=</span> <span class="mi">3</span> <span class="c1"># 1: has ball, 2: no ball, 0: background
</span>
<span class="c1"># 분류기에 사용할 입력 정보
</span><span class="n">in_features</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">roi_heads</span><span class="p">.</span><span class="n">box_predictor</span><span class="p">.</span><span class="n">cls_score</span><span class="p">.</span><span class="n">in_features</span>

<span class="c1"># 모델의 헤드 부분 교체
</span><span class="n">model</span><span class="p">.</span><span class="n">roi_heads</span><span class="p">.</span><span class="n">box_predictor</span> <span class="o">=</span> <span class="n">FastRCNNPredictor</span><span class="p">(</span><span class="n">in_features</span><span class="p">,</span> <span class="n">num_classes</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># forward 잘 되는지 테스트
</span><span class="n">a</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">,</span> <span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
<span class="n">dl</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
  <span class="n">collate_fn</span><span class="o">=</span><span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
<span class="c1"># # 학습 시
</span><span class="n">images</span><span class="p">,</span><span class="n">targets</span> <span class="o">=</span> <span class="nb">next</span><span class="p">(</span><span class="nb">iter</span><span class="p">(</span><span class="n">dl</span><span class="p">))</span>

<span class="n">output</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span><span class="n">targets</span><span class="p">)</span>
<span class="n">output</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">engine</span> <span class="kn">import</span> <span class="n">train_one_epoch</span><span class="p">,</span> <span class="n">evaluate</span>
<span class="kn">import</span> <span class="nn">utils</span>

<span class="n">device</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span> <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">cuda</span><span class="p">.</span><span class="n">is_available</span><span class="p">()</span> <span class="k">else</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cpu'</span><span class="p">)</span>

<span class="n">train_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">,</span> <span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
<span class="n">val_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">VAL_DATA_PATH</span><span class="p">,</span> <span class="n">VAL_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">train_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">train_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">num_workers</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>

<span class="n">val_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">val_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">num_workers</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>

<span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>

<span class="n">params</span> <span class="o">=</span> <span class="p">[</span><span class="n">p</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">parameters</span><span class="p">()</span> <span class="k">if</span> <span class="n">p</span><span class="p">.</span><span class="n">requires_grad</span><span class="p">]</span>
<span class="n">optimizer</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">SGD</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.005</span><span class="p">,</span>
                            <span class="n">momentum</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">weight_decay</span><span class="o">=</span><span class="mf">0.0005</span><span class="p">)</span>

<span class="c1"># lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,
#                                                 step_size=3,
#                                                 gamma=0.8)
</span><span class="n">lr_scheduler</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">lr_scheduler</span><span class="p">.</span><span class="n">MultiplicativeLR</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">lr_lambda</span><span class="o">=</span><span class="k">lambda</span> <span class="n">lr</span><span class="p">:</span> <span class="mf">0.9</span> <span class="o">**</span> <span class="n">lr</span><span class="p">)</span>

<span class="n">num_epochs</span> <span class="o">=</span> <span class="mi">15</span>

<span class="c1">#start = time.time()
</span><span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">14</span><span class="p">,</span> <span class="n">num_epochs</span><span class="o">+</span><span class="mi">1</span><span class="p">):</span>
    <span class="c1"># iteration10 마다 결과 출력
</span>    <span class="n">train_one_epoch</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">train_data_loader</span><span class="p">,</span> <span class="n">device</span><span class="p">,</span> <span class="n">epoch</span><span class="p">,</span> <span class="n">print_freq</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
    <span class="c1"># 학습률 업데이트.
</span>    <span class="n">lr_scheduler</span><span class="p">.</span><span class="n">step</span><span class="p">()</span>
    
    <span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">val_data_loader</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>

<span class="k">print</span><span class="p">(</span><span class="s">"학습 종료 "</span><span class="p">,</span> <span class="p">(</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span><span class="p">)</span> <span class="o">//</span> <span class="mi">60</span><span class="p">,</span> <span class="s">' 시간 소요'</span><span class="p">)</span>
<span class="n">torch</span><span class="p">.</span><span class="n">save</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">state_dict</span><span class="p">(),</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">WEIGHTS_PATH</span><span class="si">}</span><span class="s">faster_rcnn_</span><span class="si">{</span><span class="n">num_epochs</span><span class="si">}</span><span class="s">.pt'</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 모델 평가. 테스트용으로 하나만
</span><span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">val_dataset</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
<span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
<span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">prediction</span> <span class="o">=</span> <span class="n">model</span><span class="p">([</span><span class="n">i</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)])[</span><span class="mi">0</span><span class="p">]</span>

<span class="n">i</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">i</span><span class="p">.</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>
<span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">]):</span>
  <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">cpu</span><span class="p">(),</span> <span class="n">dtype</span> <span class="o">=</span> <span class="nb">int</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'labels'</span><span class="p">][</span><span class="n">idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
</code></pre></div></div>
<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/dfbf7705-0dd8-4cfe-9c59-12e5540a1b84" />
  </p>
  <p>
    nms 적용 전 이미지
  </p>

</div>

<p><br /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">torchvision.ops</span> <span class="kn">import</span> <span class="n">nms</span>

<span class="n">selected_idx</span> <span class="o">=</span> <span class="n">nms</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">],</span> <span class="n">prediction</span><span class="p">[</span><span class="s">'scores'</span><span class="p">],</span> <span class="n">iou_threshold</span> <span class="o">=</span> <span class="mf">0.1</span><span class="p">)</span>
<span class="n">selected_boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>
<span class="n">selected_labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'labels'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>

<span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">val_dataset</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
<span class="n">i</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">i</span><span class="p">.</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>
<span class="k">for</span> <span class="n">idx</span><span class="p">,</span><span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">selected_boxes</span><span class="p">):</span>
  <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">cpu</span><span class="p">(),</span> <span class="n">dtype</span> <span class="o">=</span> <span class="nb">int</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">selected_labels</span><span class="p">[</span><span class="n">idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/3b617a59-e9ea-4e2f-b3ff-b40096277e4a" />
  </p>
  <p>
    nms 적용 후 이미지
  </p>

</div>

<p><br /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 데이터 정의
</span><span class="n">test_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span><span class="p">,</span> <span class="n">TEST_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">test_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">test_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
<span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">test_data_loader</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
</code></pre></div></div>

<p>위의 이미지를 통해, Faster RCNN model은 일반 RCNN 모델(<a href="https://hyeonseung0103.github.io/detection/RCNN/">RCNN 포스팅 참고</a>)보다 localization을 잘 수행하고 분류 성능도 나쁘지않다는 것을 알 수 있다. Epochs를 15 정도만 했는데도 test set의 mAP50이 약 0.65였고, map@0.5:0.95는 0.407이었다. 테스트 목적이 아니라 실제로 성능을 높이기위해 에포크 수를 늘린다면 더 좋은 성능을 기록할 것이다.</p>

<p>RCNN은 selective search와 detection이 다른 네트워크에서 이루어지기때문에 한 에포크당 약 20분의 시간이 걸렸는데 Faster RCNN은은 3분 30초 정도로 RCNN보다 5배이상 빨랐다. 이를 통해 Faster RCNN은 anchor box를 기반으로 하나의 네트워크에서 region proposals, classification, box regression을 수행할 수 있기때문에 RCNN보다 훨씬 빠르면서 구현이 쉬운 모델이라는 것을 몸소 느낄 수 있었다.</p>

<h1 id="reference">Reference</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1406.4729v4.pdf">SPPNet</a></li>
  <li><a href="https://arxiv.org/pdf/1504.08083.pdf">Fast R-CNN</a></li>
  <li><a href="https://arxiv.org/pdf/1506.01497.pdf">Faster R-CNN</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[R-CNN이 다른 모델들에 비해 높은 mAP를 기록했고 이후 Fast-RCNN, Faster R-CNN, Mask R-CNN 등 여러 모델들이 R-CNN을 develop하여 object detection 분야에서 더 많은 발전을 이루어냈다. 이번 포스팅에서는 R-CNN 이후의 모델인 Fast R-CNN과 Faster R-CNN에 대해 간단하게 정리해보자.]]></summary></entry><entry><title type="html">RetinaNet 정리</title><link href="https://kimyeonz.github.io/blog/detection/RetinaNet/" rel="alternate" type="text/html" title="RetinaNet 정리" /><published>2023-11-20T00:00:00+00:00</published><updated>2023-11-20T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/RetinaNet</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/RetinaNet/"><![CDATA[<p>RetinaNet은 one stage detector의 대표주자인 YOLO, SSD보다 높은 성능을 기록하면서 Faster R-CNN보다 빠른 수행시간을 기록한 모델이다. 특히, 작은 object에 대한 detection 능력도 뛰어난데 이번 포스팅에서는 이 RetinaNet에 대해 간단히 정리해보자.</p>

<h1 id="focal-loss의-필요성">Focal Loss의 필요성</h1>
<p>Classification에서는 cross entropy를 손실함수로 많이 사용한다. 하지만, class가 imbalance할 때는 cross entropy가 상대적으로 잘 맞추고있는 class임에도 단순히 데이터 수가 많아서 loss의 많은 부분을 차지할 수 있다는 문제가 발생한다. Two stage detector에서는 RPN을 통해 객체가 있을만한 높은 확률 순으로 필터링을 수행한 후 탐지를 할 수 있지만 one stage detector에서는 모든 region(e.g. anchor box)에 대해 탐지를 수행해야하기때문에 class imbalance 문제가 더 도드라진다.</p>

<p>예를 들어, detection이 쉬운 데이터를 easy examples, 어려운 데이터를 hard examples이라고 할 때 배경과 같이 흔한 easy examples이 10,000개 자전거와 같이 예측하고자 하는 hard examples이 50개 라고 하자. 만약, Easy examples이 평균적으로 loss가 0.1이고 hard example이 1이라면 에러의 총합은 easy examples이 1,000(0.1 * 10,000), hard examples이 50(1 * 50)으로 이미 잘 맞추고있는 easy examples의 에러가 더 크게 취급된다. 결국, 우리가 잘 예측해야하는 것은 hard examples이기때문에 cross entropy를 사용하면 이와 같이 데이터의 분포는 고려되지 않은채 학습이 진행되어 학습이 불안정할 수 있다.</p>

<p>기존에는 augmentation이나 데이터셋 샘플링으로 이를 보완하려고 했지만 너무 많은 리소스가 필요하기때문에 RetinaNet에서는 Focal loss를 사용하여 이를 해결했다. 
Focal loss는 cross entropy($CE(p,y) = -\Sigma y_i ln p_i$) 공식에 가중치를 적용하는 방식이고 해당 클래스에 대한 확률이 높을수록(객체가 존재한다고 확신할수록) $\gamma$를 조절해 loss를 더 낮게하여 오히려 잘 예측하지 못한 클래스에 더 집중하도록 한다.</p>

\[FL(p_t) = -\Sigma y_i (1-p_t)^{\gamma} log(p_t)\]

<p><br /></p>
<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/088cde99-0ccb-4930-83db-e529162b5962" />
  </p>
</div>

<p><br /></p>

<p>이 Focal loss를 활용해서 Cross entropy를 손실함수로 사용했을때보다 더 좋은 정확도를 기록했다.</p>

<h1 id="feature-pyramid-networkfpn">Feature Pyramid Network(FPN)</h1>
<p>CNN에서는 층이 깊어질수록 추상적인 정보만 남아서 앞단의 세밀한 이미지 정보를 기억하기 어렵다는 문제가 있다. FPN은 이러한 문제를 해결하기 위한 기법으로 각 층의 피처맵을 예측에 사용할 피처맵과 결합하여 이미지 정보를 최대한 유지시키는 아이디어다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/862606a5-8611-40f8-b77f-cdd671cd5ddf" />
  </p>
</div>
<p><br /></p>

<p>Backbone에서 bottom-up(사이즈는 줄이고, 채널은 늘림)으로 추출한 피처맵을 top-down(사이즈를 2배로 키우고, 채널은 그대로)으로 upsampling한 피처맵과 결합하여 이 결합한 피처맵을 예측에 사용하는 것이다.
해당 피처맵에서 계산된 손실을 모두 반영하여 loss를 계산한다. 이 방법은 여러 layer의 피처맵을 예측에 사용함으로써 단일 피처맵을 사용하는 것보다 다양한 이미지 정보를 사용할 수 있다는 장점이 있다. 
또한, 각 layer의 피처맵마다 grid에 9개의 anchor box가 할당되고 anchor는 k개의 클래스 확률값과 4개의 box regression 좌표를 가진다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6c78e185-9eb2-403b-96e8-f7986da10975" />
  </p>
</div>
<p><br /></p>

<p>Faster R-CNN에 FPN을 적용했을 때 성능이 향상했고, RetinaNet의 성능이 one stage detector뿐만 아니라 two satge detector인 Faster R-CNN과 비교해도 가장 높은 것을 알 수 있다.</p>

<h1 id="구현">구현</h1>
<p>Pytorch로 RetinaNet 모델을 사용해보자(<a href="https://github.com/pytorch/vision/blob/main/torchvision/models/detection/retinanet.py">코드 참고</a>). 데이터셋 및 파일 경로 설정은 <a href="https://hyeonseung0103.github.io/detection/Fast_and_Faster_RCNN/">Fast &amp; Faster RCNN 포스팅 구현 파트 참고</a>.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 모델 정의
</span><span class="kn">import</span> <span class="nn">torchvision</span>
<span class="kn">from</span> <span class="nn">torchvision.models.detection.retinanet</span> <span class="kn">import</span> <span class="n">RetinaNetHead</span>
<span class="kn">from</span> <span class="nn">torchvision.models.detection</span> <span class="kn">import</span> <span class="n">_utils</span> <span class="k">as</span> <span class="n">det_utils</span>
<span class="kn">from</span> <span class="nn">functools</span> <span class="kn">import</span> <span class="n">partial</span>
<span class="kn">from</span> <span class="nn">torch</span> <span class="kn">import</span> <span class="n">nn</span>

<span class="n">device</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span> <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">cuda</span><span class="p">.</span><span class="n">is_available</span><span class="p">()</span> <span class="k">else</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cpu'</span><span class="p">)</span>
<span class="n">model</span> <span class="o">=</span> <span class="n">torchvision</span><span class="p">.</span><span class="n">models</span><span class="p">.</span><span class="n">detection</span><span class="p">.</span><span class="n">retinanet_resnet50_fpn_v2</span><span class="p">(</span><span class="n">pretrained</span> <span class="o">=</span> <span class="bp">True</span><span class="p">)</span>

<span class="n">num_classes</span> <span class="o">=</span> <span class="mi">3</span> <span class="c1"># has ball, no ball, background
</span>
<span class="n">in_channels</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">backbone</span><span class="p">.</span><span class="n">out_channels</span>
<span class="n">num_anchors</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">anchor_generator</span><span class="p">.</span><span class="n">num_anchors_per_location</span><span class="p">()[</span><span class="mi">0</span><span class="p">]</span>
<span class="n">norm_layer</span> <span class="o">=</span> <span class="n">partial</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">GroupNorm</span><span class="p">,</span> <span class="mi">32</span><span class="p">)</span>

<span class="n">model</span><span class="p">.</span><span class="n">head</span> <span class="o">=</span> <span class="n">RetinaNetHead</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">num_anchors</span><span class="p">,</span> <span class="n">num_classes</span><span class="p">,</span> <span class="n">norm_layer</span><span class="p">)</span>
<span class="n">model</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 데이터 정의
</span><span class="n">train_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">,</span> <span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
<span class="n">val_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">VAL_DATA_PATH</span><span class="p">,</span> <span class="n">VAL_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">train_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">train_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>

<span class="n">val_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">val_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">num_epochs</span> <span class="o">=</span> <span class="mi">30</span>
<span class="n">val_loss_tmp</span> <span class="o">=</span> <span class="mi">10000</span>
<span class="n">best_epoch_tmp</span> <span class="o">=</span> <span class="mi">1</span>
<span class="n">early_stopping_cnt</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">early_stop</span> <span class="o">=</span> <span class="mi">7</span>

<span class="n">params</span> <span class="o">=</span> <span class="p">[</span><span class="n">p</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">parameters</span><span class="p">()</span> <span class="k">if</span> <span class="n">p</span><span class="p">.</span><span class="n">requires_grad</span><span class="p">]</span>
<span class="n">optimizer</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">SGD</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span>
                            <span class="n">momentum</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">weight_decay</span><span class="o">=</span><span class="mf">0.0005</span><span class="p">)</span>

<span class="c1"># lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,
#                                                 step_size=3,
#                                                 gamma=0.9)
</span>
<span class="c1">#lr_scheduler = torch.optim.lr_scheduler.MultiplicativeLR(optimizer=optimizer, lr_lambda=lambda lr: 0.95 ** lr)
</span><span class="n">lr_scheduler</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">lr_scheduler</span><span class="p">.</span><span class="n">CosineAnnealingLR</span><span class="p">(</span><span class="n">optimizer</span><span class="p">,</span> <span class="n">T_max</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">eta_min</span><span class="o">=</span><span class="mf">0.0001</span><span class="p">)</span>

<span class="k">print</span><span class="p">(</span><span class="s">'----------------------train start--------------------------'</span><span class="p">)</span>

<span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">num_epochs</span><span class="o">+</span><span class="mi">1</span><span class="p">):</span>
  <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
  <span class="n">model</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
  <span class="n">epoch_loss</span> <span class="o">=</span> <span class="mi">0</span>
  <span class="n">prog_bar</span> <span class="o">=</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">train_data_loader</span><span class="p">,</span> <span class="n">total</span><span class="o">=</span><span class="nb">len</span><span class="p">(</span><span class="n">train_data_loader</span><span class="p">))</span>

  <span class="k">for</span> <span class="n">images</span><span class="p">,</span> <span class="n">targets</span> <span class="ow">in</span> <span class="n">prog_bar</span><span class="p">:</span>
    <span class="n">images</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">image</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">)</span>
    <span class="n">targets</span> <span class="o">=</span> <span class="p">[{</span><span class="n">k</span><span class="p">:</span> <span class="n">v</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">t</span><span class="p">.</span><span class="n">items</span><span class="p">()}</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">targets</span><span class="p">]</span>

    <span class="n">loss_dict</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span> <span class="n">targets</span><span class="p">)</span>

    <span class="n">optimizer</span><span class="p">.</span><span class="n">zero_grad</span><span class="p">()</span>
    <span class="n">loss</span> <span class="o">=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">loss</span> <span class="k">for</span> <span class="n">loss</span> <span class="ow">in</span> <span class="n">loss_dict</span><span class="p">.</span><span class="n">values</span><span class="p">())</span>
    <span class="n">loss</span><span class="p">.</span><span class="n">backward</span><span class="p">()</span>
    <span class="n">optimizer</span><span class="p">.</span><span class="n">step</span><span class="p">()</span>
    <span class="n">epoch_loss</span> <span class="o">+=</span> <span class="n">loss</span><span class="p">.</span><span class="n">item</span><span class="p">()</span>
  <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'epoch : </span><span class="si">{</span><span class="n">epoch</span><span class="si">}</span><span class="s">, Loss : </span><span class="si">{</span><span class="n">epoch_loss</span><span class="si">}</span><span class="s">, time : </span><span class="si">{</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>

  <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">epoch_val_loss</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="n">val_start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
    <span class="k">for</span> <span class="n">images</span><span class="p">,</span> <span class="n">targets</span> <span class="ow">in</span> <span class="n">val_data_loader</span><span class="p">:</span>
        <span class="n">images</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">image</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">)</span>
        <span class="n">targets</span> <span class="o">=</span> <span class="p">[{</span><span class="n">k</span><span class="p">:</span> <span class="n">v</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">t</span><span class="p">.</span><span class="n">items</span><span class="p">()}</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">targets</span><span class="p">]</span>

        <span class="n">val_loss_dict</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span> <span class="n">targets</span><span class="p">)</span>
        <span class="n">epoch_val_loss</span> <span class="o">+=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">loss</span> <span class="k">for</span> <span class="n">loss</span> <span class="ow">in</span> <span class="n">val_loss_dict</span><span class="p">.</span><span class="n">values</span><span class="p">())</span>

    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'Val Loss : </span><span class="si">{</span><span class="n">epoch_val_loss</span><span class="si">}</span><span class="s">, time : </span><span class="si">{</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">val_start</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>
    <span class="k">if</span> <span class="n">epoch_val_loss</span> <span class="o">&lt;</span> <span class="n">val_loss_tmp</span><span class="p">:</span>
        <span class="n">early_stopping_cnt</span> <span class="o">=</span> <span class="mi">0</span>
        <span class="n">best_epoch_tmp</span> <span class="o">=</span> <span class="n">epoch</span>
        <span class="n">val_loss_tmp</span> <span class="o">=</span> <span class="n">epoch_val_loss</span>
        <span class="n">torch</span><span class="p">.</span><span class="n">save</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">state_dict</span><span class="p">(),</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">WEIGHTS_PATH</span><span class="si">}</span><span class="s">retinanet_</span><span class="si">{</span><span class="n">num_epochs</span><span class="si">}</span><span class="s">.pt'</span><span class="p">)</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="n">early_stopping_cnt</span> <span class="o">+=</span> <span class="mi">1</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'현재까지 best 모델은 Epochs </span><span class="si">{</span><span class="n">best_epoch_tmp</span><span class="si">}</span><span class="s">번째 모델입니다.'</span><span class="p">)</span>

  <span class="k">if</span> <span class="n">early_stopping_cnt</span> <span class="o">==</span> <span class="n">early_stop</span><span class="p">:</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">early_stop</span><span class="si">}</span><span class="s">번 동안 validation 성능 개선이 없어 학습을 조기 종료합니다.'</span><span class="p">)</span>
    <span class="k">break</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 데이터 정의
</span><span class="n">test_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span><span class="p">,</span> <span class="n">TEST_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">test_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">test_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">test_data_loader</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span> <span class="c1"># mAP@0.5:0.95 0.635, mAP@0.5 0.893
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">torchvision.ops</span> <span class="kn">import</span> <span class="n">nms</span>
<span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">test_dataset</span><span class="p">[</span><span class="mi">10</span><span class="p">]</span>
<span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
<span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">prediction</span> <span class="o">=</span> <span class="n">model</span><span class="p">([</span><span class="n">i</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)])[</span><span class="mi">0</span><span class="p">]</span>

<span class="n">selected_idx</span> <span class="o">=</span> <span class="n">nms</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">],</span> <span class="n">prediction</span><span class="p">[</span><span class="s">'scores'</span><span class="p">],</span> <span class="n">iou_threshold</span> <span class="o">=</span> <span class="mf">0.5</span><span class="p">)</span>
<span class="n">selected_boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>
<span class="n">selected_labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'labels'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>
<span class="n">selected_scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'scores'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>

<span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">test_dataset</span><span class="p">[</span><span class="mi">10</span><span class="p">]</span>
<span class="n">i</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">i</span><span class="p">.</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>
<span class="k">for</span> <span class="n">idx</span><span class="p">,</span><span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">selected_boxes</span><span class="p">):</span>
  <span class="k">if</span> <span class="n">selected_scores</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span> <span class="o">&gt;</span> <span class="mf">0.9</span><span class="p">:</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">cpu</span><span class="p">(),</span> <span class="n">dtype</span> <span class="o">=</span> <span class="nb">int</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">selected_labels</span><span class="p">[</span><span class="n">idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/3c307ee2-bff6-44d7-8df5-05ee509e411e" />
  </p>
  <p>nms 적용 후 RetinaNet 모델의 test 이미지 결과(confidence threshold 0.9)</p>
</div>
<p><br /></p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/ab693c59-3161-4f15-a572-a1b4a0ecf4d1" />
  </p>
  <p>nms 적용 후 SSDLite 모델의 test 이미지 결과(confidence threshold 0.9)</p>
</div>
<p><br /></p>

<p>위의 이미지를 보면, RetinaNet은 confidence score가 0.9이상인 박스들을 추출하면 원하는 객체를 올바르게 탐지하는 반면 SSD는 5개의 객체 중 3개의 객체만 탐지한다. 또한, RetinaNet은 해당 객체들 중 볼을 소유하고 있는 객체의 클래스를 1이라고 올바르게 예측했다.</p>

<p>RetinaNet은 Focal Loss와 FPN을 활용하여 당시 one stage detector에서 좋은 성능을 보였던 YOLO와 SSD보다 뛰어난 정확도를 가졌고, two stage detector인 Faster R-CNN보다도 높은 정확도를 기록했다고 논문에 언급되었다.</p>

<p>실제로 구현을 해보니 RetinaNet은 SSD 0.546, Faster R-CNN 0.407, YOLOv1 0.34보다 높은 0.635의 mAP를 기록했다. 비록 YOLO와 SSD보다는 학습 시간(custom dataset 기준 한 에포크당 20초)이 느리긴하지만 한 에포크당 학습 시간이 Faster R-CNN(한 에포크당 3분 30초)보다 약 1.4배 정도 더 빠른 2분 30초의 시간이 소요됐다.</p>

<h1 id="reference">Reference</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1708.02002.pdf">RetinaNet Paper</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[RetinaNet은 one stage detector의 대표주자인 YOLO, SSD보다 높은 성능을 기록하면서 Faster R-CNN보다 빠른 수행시간을 기록한 모델이다. 특히, 작은 object에 대한 detection 능력도 뛰어난데 이번 포스팅에서는 이 RetinaNet에 대해 간단히 정리해보자.]]></summary></entry><entry><title type="html">Mask R-CNN 논문 요약</title><link href="https://kimyeonz.github.io/blog/segmentation/Mask_RCNN/" rel="alternate" type="text/html" title="Mask R-CNN 논문 요약" /><published>2023-10-19T00:00:00+00:00</published><updated>2023-10-19T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/segmentation/Mask_RCNN</id><content type="html" xml:base="https://kimyeonz.github.io/blog/segmentation/Mask_RCNN/"><![CDATA[<h1 id="abstract">Abstract</h1>
<p>본 논문에서는 object instance segmentation task에 대한 simple, flexible, general한 framework에 대해 소개한다. 본 연구의 접근방식은 효과적으로 객체를 탐지하는 동시에 각각의 인스턴스에 대해 높은 수준의 segmentation mask를 생성한다. Mask R-CNN은 Faster R-CNN에 layer를 확장시켜 object mask를 병렬로 예측하고 기존의 box recognition도 그대로 수행할 수 있다. Mask R-CNN은 학습이 쉽고 Faster R-CNN보다 연산이 크게 추가되지 않은 5 fps의 속도로 동작한다. 나아가, 사람의 동작을 예측하는 task 등 다양한 task를 수행할 수 있다. COCO datset에서 instance segmentation, bounding box object detection, person keypoint detection에 대해 좋은 성적을 거두었다. COCO 2016의 우승팀을 포함하여 기존에 존재한 모든 single model보다 우수한 성적을 거두었다. 본 연구가
향후 instance-level recognition 분야에 큰 도움이 되길 바란다.</p>

<p><br /><br /></p>

<h1 id="details">Details</h1>
<h2 id="introduction">Introduction</h2>
<ul>
  <li>Object detection과 semantic segmentation은 Fast/Faster R-CNN과 FCN처럼 좋은 baseline을 사용해 짧은 기간동안 많은 발전을 이뤘다.</li>
  <li>본 연구의 목표는 다른 task들과 비슷한 수준으로 instance segmentation이 가능한 framework를 만드는 것이다.</li>
  <li>Instance segmentation은 모든 객체에 대해 정확한 detection과 segmentation이 필요하기때문에 어렵다. 따라서, 복잡한 방법이 좋은 결과를 불러올 것이라고 생각했는데 굉장히 단순하고 유연하고 빠른 시스템이 instance segmentation의 SOTA를 넘었다.</li>
  <li>Mask R-CNN은 Faster R-CNN을 확장시켜 각 RoI에 대해 segmentation mask를 예측하고, 병렬적으로 기존과 같이 classification, bounding box regression을 수행한다. Mask branch는 작은 FCN이 RoI에 적용된 것이고 pixel 단위로 segmentation mask를 예측한다. 또한, mask branch는 작은 계산만 추가되기때문에 기존처럼 빠른 시스템을 만들 수 있었다.</li>
  <li>Faster R-CNN은 pixel-to-pixel로 설계되지 않았기때문에 <strong>RoIPool</strong>이 특징 추출을 위해 공간적인 정보를 coarse(거칠게. pooling을 수행하여 일부 공간정보가 손실되기때문에 이렇게 어느 정도 희생을 감소해서 특징을 추출하는 것을 coarse하다고 함)하게 추출한다. 이러한 misalignment 문제를 해결하기 위해 본 연구에서는 간단하고 공간적인 제약에서 자유로운 <strong>RoIAlign</strong> 기법을 사용해서 공간 위치를 잘 보존한다. 이 간단한 변화가 mask 정확도를 10%에서 50%까지 올렸다.</li>
  <li>또한, mask와 class prediction을 분리하는 것이 매우 중요하다는 것을 알았다. 각각의 class끼리 영향을 주지않도록 class에 대해 binary mask를 독립적으로 수행했고, class를 예측하기위해서는 RoI classification에만 의존했다.</li>
  <li>반면, FCN은 보통 픽셀단위에서 multi-class categorization을 수행하는데 이는 segmentation과 classification을 묶은 것이고, 이 방법은 경험상 instance segmentation 분야에서 좋지 않은 방법이다.</li>
  <li>Mask R-CNN은 종과 호루라기 말고는 COCO instance segmentation task에서 이전 대회의 SOTA를 뛰어넘었다. GPU에서 프레임당 200ms가 소요됐고 COCO dataset을 학습했을 때 하루에서 이틀 정도의 시간이 걸렸다. 이렇게 학습과 test가 빠르고 정확도도 뛰어나기때문에 향후 연구에 잘 쓰일 모델이라고 기대한다.</li>
  <li>Human pose estimation에서도 좋은 속도와 성능을 기록했다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/7db88c37-8c1c-4f7a-85a0-371ff7049cc1" />
  </p>
</div>

<p><br /></p>

<p><br /><br /></p>

<h2 id="related-work">Related Work</h2>
<p><strong>R-CNN</strong></p>
<ul>
  <li>R-CNN은 빠르고 정확한 모델로 개선시키기위해 RoIPool을 사용하여 피처맵에서 미리 정의된 RoI에 대한 정보를 추출할 수 있도록했다.</li>
  <li>Faster R-CNN은 RPN을 사용해서 이 기술을 더 발전시켰고 다양한 방법을 동원해서 모델을 유연하고 강건하게 개선시켰다. 그리고 현재, object detection 분야를 이끌어나가는 모델이 되었다.</li>
</ul>

<p><strong>Instance Segmentation</strong></p>
<ul>
  <li>R-CNN의 영향으로 instance segmentation에 대한 많은 접근들이 segment proposals을 기반으로 이루어져있다.</li>
  <li>DeepMask와 이후의 연구들에서는 segment 후보를 제안하도록 학습한 후 Fast R-CNN을 통해 분류를 수행한다. 이 방법은 segmentation precedes(선행, 우선) recognition으로 느리고 정확도가 낮다는 단점이 있다.</li>
  <li>본 연구에서는 mask와 class label에 대해 병렬로 예측을 수행해서 간단하고 유연한 모델을 만들었다.</li>
  <li>가장 최근 Li의 연구에서는 segment proposal system과 object detection system을 결합한 fully convolutional instance segmentation(FCIS)를 개발했다. 이 모델은 position sensitive output channels 집합을 fully convolutional하게 예측하는 것으로 이 채널이 object classes, boxes, masks를 동시에 처리해서 모델을 빠르게 한다.</li>
  <li>하지만, FCIS는 overlapping instance에 의한 에러와 가짜 edges를 만들어 instance segmentation을 수행하기에는 어려움을 겪었다.</li>
  <li>다른 접근으로는, 좋은 semantic segmentation을 기반으로 instance segmentation을 수행하는 것이다. FCN outputs과 같이 pixel 단위로 classification을 수행하고 이 결과를 가지고 같은 카테고리에 있더라도 다른 instance로 잘라내려고 시도한다.</li>
  <li>Mask R-CNN은 이러한 segmentation first 기법과 다르게 instance first 기법을 사용한다. 그리고 앞으로 이 두가지 기법을 더 깊게 통합하는 연구가 향후에 있을 것으로 예상한다.</li>
</ul>

<p><br /><br /></p>

<h2 id="mask-r-cnn">Mask R-CNN</h2>
<p>Mask R-CNN은 개념적으로 간단하다. Faster R-CNN은 한 객체에 대해 class label과 bounding box offsets 이 2가지 outputs을 내는데 Mask R-CNN은 여기에 branch를 하나 더 추가해서 object mask까지 outputs으로 내는 것이다. 그럼 지금부터는 Fast/Faster R-CNN이 놓쳤고, pixel-to-pixel alignment를 포함한 Mask R-CNN의 핵심 요소에 대해 알아보자.</p>

<p><br /></p>

<p><strong>Faster R-CNN &amp; Mask R-CNN</strong></p>
<ul>
  <li>Faster R-CNN은 1 stage는 RPN, 2 stage는 RoIPool을 통해 특징을 추출하고 분류 및 box regression을 수행하는 Fast R-CNN으로 이루어져있다. Mask R-CNN도 2 stage 모델로 1 stage는 RPN, 2 stage는 class/box offsets/binary mask tasks를 각 RoI에 대해 병렬적으로 수행한다. 이것은 classification에 의존한 mask predictions을 사용하는 최근 모델들과는 다르다는 것을 보여준다.</li>
  <li>학습 과정에서는 RoI에 대해 $L = L_{cls} + L_{box} + L_{mask}$로 multi-task loss를 정의한다. Classification loss와 box loss는 Fast R-CNN과 같고 mask branch는 mask resolution이 m x m, 총 class 수가 $K$라고 했을 때 각 RoI에 대해 $Km^2$의 사이즈를 갖는 output을 도출한다. 이를 적용하기 위해 pixel 단위로 sigmoid를 사용했고 $L_{mask}$를 average binary cross-entropy loss로 정의했다. RoI내에서 mask들끼리는 서로 연관되어있지않고 독립적인 loss로 사용된다.</li>
  <li>이렇게 Mask R-CNN은 masked class에 대해서 전용 분류기를 따로 사용해서 mask와 classification 예측을 분리하는 구조다(decouples). FCN에서는 semantic segmentation을 수행할 때 픽셀단위로 softmax와 multinomial cross-entropy loss를 사용해서 픽셀마다 mask들이 다른 class들과 연관되어 mask, classification 예측이 함께 이루어지지만, Mask R-CNN에서는 픽셀 단위로 sigmoid와 binary loss를 사용하기때문에 두 가지 예측에서 class들끼리 엮여지않고 각 예측이 독립적으로 수행되어 연산량의 측면에서 훨씬 효율적이다.</li>
</ul>

<p><br /></p>

<p><strong>Mask Representation</strong></p>
<ul>
  <li>Mask는 class labels과 box offsets과 달리 fc층에 의해 어쩔 수 없이 공간정보가 손실된 짧은 벡터로 변환되는데 mask에 대한 공간적인 정보를 추출하는 것은 CNN에서 pixel-to-pixel로 작업이 이루어지기때문에 크게 걱정할 필요없다. 특히, 각 RoI에 대해 mxm으로 mask를 예측할 때 FCN(Fully Convolutional Layer)을 사용하기때문에 공간정보의 손실이 존재하는 벡터가 아닌 상태로 mxm의 공간정보를 잘 유지시킬 수 있다.</li>
  <li>이런 pixel-to-pixel 방법을 활용하기 위해서는 그 자체가 작은 피처맵이라고 할 수 있는 RoI의 픽셀당 공간 정보가 잘 정렬이 되어있어야하는데 본 연구팀의 RoIAlign 기법이 mask prediction을 위한 핵심 역할을 수행한다.</li>
</ul>

<p><br /></p>

<p><strong>RoIAlign</strong></p>
<ul>
  <li>RoIPool은 다양한 크기의 RoI를 CNN에 집어넣기위해 고정된 피처맵(e.g. 7x7)으로 변환하는 역할을 수행했다. 먼저, 소수 형태로 이루어져있는 RoI를 이산화(정수형태로 변환. e.g. [x/16] x를 16으로 나누고 정수부분만 사용)하고 해당 RoI를 고정된 크기의 피처맵으로 출력하기 위해 주로 max pooling을 사용하여 aggregate 시킨다. 하지만, 이렇게 소수 형태로 이루어져있는 RoI와 quantization(연속적인 좌표값을 피처 맵의 그리드에 딱 맞게 떨어지는 값으로 변환하는 과정)이 수행된 피처맵이 misalignment되는 경우가 많기때문에 classification에는 영향을 미치지않더라도 pixel 단위로 이루어지는 segmentation과 같은 task에는 안 좋은 영향을 끼친다.</li>
  <li>RoIAlign은 이러한 문제를 해결하기 위해 RoIPool에서 harsh quantization된 피처를 재조정해서(실제로 재조정하는게 아니고 개념적으로 보면 더 정확한 align이 이루어짐) input과 피처맵을 정렬시킨다. 방법은 이산화를 시켰던 방법을 소수부분까지 그대로 사용해서 quantization을 막고(e.g. [x/16] -&gt; x/16)bilinear interpolation을 통해 각 RoI의 픽셀들을 가중합시켜서 기존 RoI와 유사하도록 피처맵을 구성한다.</li>
  <li>Figure 3.은 점선이 피처맵, 실선이 RoI(2x2 bin을 가진)라고 할 때, 더 정확한 RoI 피처맵을 만들기위해 각각의 bin을 4개의 sampling points를 사용해서 또 4등분 하여 subsample을 만들고 이 subsample에 대해 bilinear interpolation을 적용하고 최종적으로 max pooling이나 average pooling을 적용해 특징을 추출하는 것을 설명하는 그림이다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/cb14b555-387b-4b9a-b1b0-7f15424ccb43" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>아래 그림은 subsampling에서 bilinear interpolation을 사용하고 최종적으로 max pooling을 적용하는 예시이다. Bilinear interpolation을 보면 subsample에서 다른 pixel들이 포함된 만큼만 공평하게 비율을 가져가는 즉, 적절한 가중합이 이루어진 것을 알 수 있다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/34b47475-3713-4187-a760-5c8f715d5c1e" />
  </p>
</div>

<p><br /></p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/cc3f305a-3e7f-4b7f-b35a-58379229a7c0" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>RoIAlign을 통해 큰 성능 향상을 이루어냈다.</li>
</ul>

<p><br /></p>

<p><strong>Network Architetecture</strong></p>
<ul>
  <li>Mask R-CNN의 일반화 능력을 입증하기위해 다양한 architecture를 조합하여 사용했다.
    <ul>
      <li>전체 이미지에 대해 특징을 추출한 convolutional backbone architecture</li>
      <li>network head 부분에서는 bounding box recognition(classification &amp; regression)</li>
      <li>각 RoI에 대해 mask prediction</li>
    </ul>
  </li>
  <li>Backbone
    <ul>
      <li>ResNet과 ResNeXt의 50 혹은 101 layers를 backbone으로 사용했는데 만약 ResNet-50의 final convolutional layer에서 C4라고 불리는 4번째 stage까지의 feature를 사용했으면 ResNet-50-C4라고 칭한다.</li>
      <li>ResNet-FPN(Feature Pyramid Network) backbone을 사용했을 때 정확도와 속도면에서 우수했다.</li>
    </ul>
  </li>
  <li>Network head
    <ul>
      <li>Head에는 Faster R-CNN에 mask prediction branch가 있는 FCN을 추가한게 전부다.</li>
      <li>구조는 Figure 4.를 참고하자.</li>
    </ul>
  </li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/43f9c045-801b-453e-b50d-b0e06c83f4fc" />
  </p>
</div>

<p><br /></p>

<h3 id="1-implementation-details">1. Implementation Details</h3>
<p><strong>Training</strong></p>
<ul>
  <li>Fast/Faster R-CNN과 동일한 하이퍼파라미터를 사용했고 이는 object detection을 위해서 사용된 것이지만 instance segmentation에도 잘 맞았다.</li>
  <li>Fast R-CNN에서 IoU thresholds를 0.5로 사용했다. Mask R-CNN에서 mask loss $L_{mask}$는 positive RoI에 대해서만 정의된다.</li>
  <li>전처리에는 image resizing, mini-batch 2 images per GPU, N RoIs per image, 1:3 positive:negative ratio가 적용됐다. N은 C4 backbone에서는 64, FPN에서는 512로 적용했다.</li>
  <li>8 GPUs(mini batch 2 * 8 = 16), 160k iterations, lr 0.02(120k iteration에 10배 감소), weight decay 0.0001, momentum 0.09.</li>
  <li>RPN anchors box는 5 scales과 3 aspect ratios를 가졌고 RPN은 명시하지않는 한 Mask R-CNN과 별개로 학습되어 features를 공유하지 않지만, 본 연구에서는 두 네트워크가 backbone이 같기때문에 공유할 수 있다.</li>
</ul>

<p><br /></p>

<p><strong>Inference</strong></p>
<ul>
  <li>Test에는 RoI가 C4에서는 300, FPN에서는 1000이 사용된다. NMS가 적용된 Mask branch는 가장 성능이 좋은 100 detection boxes를 예측한다. 이 방법은 학습과 다른 parallel computation이지만 더 갯수가 적고 정확한 RoI를 사용한 것이기때문에 inference 속도를 높이고 정확도를 증가시킬 수 있다.</li>
  <li>Mask branch는 RoI당 K(class 갯수)개의 masks를 예측할 수 있지만 모든 class에서의 mask가 아니라 classification branch에서 예측한 k라는 class에 대해서만 mask 예측을 수행한다(학습에서는 모전체 class에 대한 mask였지만 inference에서는 굳이 전체 class에 대해서 할 필요없음).</li>
  <li>m x m 사이즈로 소수로 이루어진 mask output은 RoI와 같은 크기로 resizing되고 threshold 0.5를 기준으로 binarized 된다.</li>
  <li>Top 100 detection box로만 mask를 예측했기때문에 Faster R-CNN에서 비해 overhead가 조금밖에 증가하지않았다.</li>
</ul>

<p><br /><br /></p>

<h2 id="experiments-instance-segmentation">Experiments: Instance Segmentation</h2>
<h3 id="1-main-results">1. Main Results</h3>
<p><br /></p>
<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/bf920c64-226b-43e7-b545-c950e21a5762" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Mask R-CNN과 SOTA 모델을 비교했고, COCO datasets에서 다양한 AP scale을 가지고 실험을 진행했다. AP는 mask IoU로 평가했다.</li>
  <li>Table 1.을 보면 Mask R-CNN은 모든 AP scale에서 이전 SOTA 모델들을 뛰어넘었다. 종과 호루라기 객체말고는 RestNet-101-FPN backbone모델이 multi-scale train/test, horizontal flip test, online hard example mining등 다양한 기법이 적용된 FCIS+++을 뛰어넘었다.</li>
  <li>Figure 6.는 Mask R-CNN과 FCIS+++을 비교한 것으로 FCIS+++은 겹치는 인스턴스에 대해 artifacts(가짜 인스턴스)를 만들어내는 반면 Mask R-CNN은 인스턴스가 겹쳐도 좋은 예측을 수행하고 있다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="1000" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/2ca46f6a-bb37-487f-bfae-48d5dd174ab2" />
  </p>
</div>

<p><br /></p>

<h3 id="2-ablation-experiments">2. Ablation Experiments</h3>
<p><br /></p>
<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/823421e9-5175-4340-a540-ca62c02da037" />
  </p>
</div>

<p><br /></p>

<p><strong>Architecture</strong></p>
<ul>
  <li>Mask R-CNN은 다양한 backbones을 사용했고 깊은 네트워크(50 vs 101)와 FPN, ResNeXt과 같은 advanced 모델들이 성능 향상에 도움이 됐다(Table 2a). 하지만, 모든 프레임워크가 자동으로 깊고 advanved한 네트워크에서 이점을 누릴 수 있는 것은 아니다.</li>
</ul>

<p><br /></p>

<p><strong>Multinomial vs Independent Masks</strong></p>
<ul>
  <li>Table 2b는 FCN처럼 pixel 단위로 softmax와 multinomial loss를 사용하느냐 Mask R-CNN처럼 sigmoid와 binary loss를 사용하느냐의 결과를 비교한 것이다. 쉽게 말해서, classificaiton과 mask가 couple이나 decouple이냐의 비교다.</li>
  <li>Couples tasks를 사용했을때 AP가 5.5 points 하락했다. 이를 통해 다른 class와 상관없이 binary mask를 사용했을때 모델이 더 잘 학습한다는 것을 알 수 있다.</li>
</ul>

<p><br /></p>

<p><strong>Class-Specific vs Class-Agnostic Masks</strong></p>
<ul>
  <li>본 연구에서는 기본적으로 class-specific masks 즉, 클래스마다 mxm의 mask를 만들어냈다. 흥미롭게도 Mask R-CNN에 클래스와 상관없이 하나의 mxm mask만을 생성하는 class-agnostic masks를 적용했더니 생각보다는 효과적이었다. Class annostic masks는 29.7 AP, specific masks(ResNet-50-C4)는 30.3 masks로 큰 차이가 없었다.</li>
  <li>하지만 기본적으로 사용한 class-specific masks가 성능이 더 좋기때문에 mask와 classification간의 decouple이 중요하다는 것을 보여준다.</li>
</ul>

<p><br /></p>

<p><strong>RoIAlign</strong></p>
<ul>
  <li>RoIAlign의 효과를 검증할 때는 ResNet-50-C4 backbone(stride 16)을 사용했고 RoIAlign은 RoIPool보다 3 points 높은 AP_75를 기록했다(Tabel 2c).</li>
  <li>MNC의 RoIWarp(quantization이 되는 것을 허용하지만, bilinear interpolation은 적용하는 기법)와도 비교해본 결과 RoIWarp는 RoIPool보다는 좋지만 RoIAlign보다는 낮은 AP를 기록한다.</li>
  <li>추가적으로 ResNet-50-C5 backbone(stride 32)을 설계해 RoIAlign을 비교해보았는데 결과적으로 stride가 더 큰 C5 backbone 모델의 AP가 C4보다 높았다(Tabel 2c, 2d). 아무래도 stride가 커져서 RoIPool의 misalignment가 더 심해졌을것이고 이에 따라 RoIAlign의 효과는 커졌을 것으로 판단된다. RoIAlign은 stride가 클수록 detection과 segmentation을 더 잘 수행할 수 있게한다는 것을 알 수 있다.</li>
  <li>FPN을 backbone으로 사용했을 때도 RoIAlign의 성능이 더 좋았다(Table 6.).</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/8d2ebef7-f077-402e-ae22-0e9435446e89" />
  </p>
</div>

<p><br /></p>

<p><strong>Mask Branch</strong></p>
<ul>
  <li>Segmentation은 FCN을 이용한 pixel-to-pixel task를 수행한다. FCN의 효과를 확인하기위해 FCN과 MLP(multi layer perceptions)를 비교했다(ResNet-59-FPN backbone 사용, Table 2e).</li>
  <li>FCN을 사용했을 때 MLP보다 AP가 더 높았다. ResNet-50-FPN backbone을 사용하면 head부분은 FCN으로 이루어져있기때문에 MLP와의 공정한 비교를 위해 FCN head부분은 pretrain 시키지 않았다.</li>
</ul>

<p><br /></p>

<h3 id="3-bounding-box-detection-results">3. Bounding Box Detection Results</h3>
<ul>
  <li>객체 탐지 성능도 확인하기 위해 COCO datasets에서 SOTA 모델들과 비교했다(mask output을 제외하고 classification과 box outputs만 사용). ResNet-101-FPN backbone을 사용하면 이전 여러 SOTA 모델들보다 성능이 더 좋았다.</li>
  <li>다른 비교를 진행하기 위해 Mask R-CNN에서 mask branch를 제거한 <strong>Faster R-CNN, RoIAlign</strong> 모델을 만들었다. 이 모델은 RoIAlign을 통해 FPN보다 성능이 더 좋았고 box AP는 Mask R-CNN보다 성능이 좋지 않았다. 이를 통해 Mask R-CNN이 multi-task(detection, clssification, segementation)에 효과적이라는 것을 알 수 있다.</li>
  <li>Mask와 box AP의 성능차이가 크지않는 것을 보아 Mask R-CNN은 object detection과 더 어려운 instance segmentation task의 갭을 크게 줄였다고 할 수 있다(Table 1, 3 비교).</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6a6ef5b0-52fa-4127-b0c0-7adeb18cd094" />
  </p>
</div>

<p><br /></p>

<h3 id="4-timing">4. Timing</h3>
<p><strong>Inference &amp; Training</strong></p>
<ul>
  <li>RPN과 Mask R-CNN stages가 features share하도록 ResNet-101-FPN 모델을 사용했을 때 이미지당 195ms의 처리시간이 소요됐다. ResNet-101-C4는 무겁기때문에 더 오랜 시간이 소요돼서 추천하지 않는다.</li>
  <li>Mask R-CNN이 빠르긴하지만 속도를 최적화시킨 것이 아니기때문에 다양한 이미지 사이즈, proposal number 등을 사용해서 향후 연구에서 정확도와 속도를 더 향상시킬 것으로 기대한다.</li>
  <li>Mask R-CNN은 학습도 빠르다. ResNet-50-FPN에서 8 GPU, COCO trainval_35k를 사용했을 때 32시간, RestNet-101-FPN은 44시간이 걸렸다.</li>
</ul>

<p><br /><br /></p>

<h2 id="mask-r-cnn-for-human-pose-estimation">Mask R-CNN for Human Pose Estimation</h2>
<p>Mask R-CNN으로 K개의 key point types(e.g., left shoulder, right elbox..)을 masking하도록 했다. 이 task가 Mask R-CNN이 flexible하다는 것을 보여준다. 본 연구팀은 human pose 도메인에 대한 지식이 많지 않기때문에 이 task는 Mask R-CNN의 generality를 검증하는 용도로 사용한다. Mask R-CNN에 도메인 지식까지 겸비되면 더 좋은 모델이 만들어질 것이다.</p>

<p><br /></p>

<p><strong>Implementation Deatils</strong></p>
<ul>
  <li>새로운 tasks를 수행하기위해 한 instance의 K개의 key points에 대해 training target을 mxm mask 중 하나의 pixel만 foreground(key points)로 라벨링된 one-hot mxm binary mask을 예측하는 것으로 만들었다. Instance segmentation처럼 K keypoints로 서로 독립적으로 예측된다.</li>
  <li>ResNet-FPN을 조금 변형시켜서 사용했고 keypoint-level localization에서는 mask task에 비해 상대적으로 high resolution output이 필요하다는 것을 알았다.</li>
</ul>

<p><br /></p>

<p><strong>Main Resluts and Ablations</strong></p>

<p><br /></p>
<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/8c49de3b-fb75-4e38-9199-afc236f5e8d5" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>ResNet-50-FPN, Key point에 대한 AP를 평가지표로 사용했을 때 2016 COCO keypoint detection winner보다 높은 성능을 기록했다. Mask R-CNN이 더 간단하고 빠르다.</li>
  <li>더 중요한 것은 predict boxes, segments, keypoints을 동시에 수행하는 unified 모델인데도 5 fps를 기록했다.</li>
  <li>Key point detection task는 어쩌면 mask보다 더 정확한 localization이 필요하기때문에 RoIAlign의 중요성은 역시나 RoIPool보다 컸다.</li>
</ul>

<p><br /><br /></p>

<h1 id="개인적인-생각">개인적인 생각</h1>
<ul>
  <li>발전이 더딘 instance segmentation분야에서 큰 개선을 이뤄냈고 6년이 지난 지금도 instance segmentation 분야를 대표하는 모델이라는 점에서 굉장히 의미있는 연구인 것 같다.</li>
  <li>RoIPool이 가지고 있었던 misalignment 문제를 RoIAlign이라는 간단한 기법으로 해결했다는 점이 인상깊었다.</li>
  <li>Human pose estimation이라는 기존과는 전혀 다른, 심지어 스스로도 지식이 부족하다고 언급한 도메인에 도전해서 일반화 성능을 검증하려고 한 연구팀의 열정을 느낄 수 있었다.</li>
  <li>Classification과 mask prediction을 분리해서 좋은 성과를 냈지만 decouple의 이유로 class 갯수만큼 mask를 만들어야하고, binary가 softmax보다는 연산이 간단하겠지만 결국 불필요하게 모든 class에 대해 mask를 생성해야한다는 점에서 의문이 들었다. 향후 연구에서는 어떤 방법으로 mask tasks를 좀 더 간단하게 만들지 궁금하다.</li>
  <li>Mask R-CNN도 결국 RPN과 분류기로 이루어진 2 stage 모델이기때문에 1 stage 모델보다 속도면에 있어서는 불리하다. 논문의 저자가 언급한 것처럼 속도에 최적화시킨 모델이 아니기떄문에 Mask R-CNN을 기반으로 향후 연구에서 어떻게 빠르고 정확한 instance segmentation model을 만들지 기대된다.</li>
</ul>

<p><br /><br /></p>

<h1 id="이미지-출처">이미지 출처</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1703.06870.pdf">Mask R-CNN Paper</a></li>
  <li><a href="https://blog.kubwa.co.kr/%EB%85%BC%EB%AC%B8%EB%A6%AC%EB%B7%B0-mask-r-cnn-2018-mask-r-cnn-%EC%8B%A4%EC%8A%B5-w-pytorch-cd525ea9e157">RoI Align Image</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Segmentation" /><summary type="html"><![CDATA[Abstract 본 논문에서는 object instance segmentation task에 대한 simple, flexible, general한 framework에 대해 소개한다. 본 연구의 접근방식은 효과적으로 객체를 탐지하는 동시에 각각의 인스턴스에 대해 높은 수준의 segmentation mask를 생성한다. Mask R-CNN은 Faster R-CNN에 layer를 확장시켜 object mask를 병렬로 예측하고 기존의 box recognition도 그대로 수행할 수 있다. Mask R-CNN은 학습이 쉽고 Faster R-CNN보다 연산이 크게 추가되지 않은 5 fps의 속도로 동작한다. 나아가, 사람의 동작을 예측하는 task 등 다양한 task를 수행할 수 있다. COCO datset에서 instance segmentation, bounding box object detection, person keypoint detection에 대해 좋은 성적을 거두었다. COCO 2016의 우승팀을 포함하여 기존에 존재한 모든 single model보다 우수한 성적을 거두었다. 본 연구가 향후 instance-level recognition 분야에 큰 도움이 되길 바란다.]]></summary></entry><entry><title type="html">SSD: Single Shot MultiBox Detector 논문 요약</title><link href="https://kimyeonz.github.io/blog/detection/SSD/" rel="alternate" type="text/html" title="SSD: Single Shot MultiBox Detector 논문 요약" /><published>2023-10-17T00:00:00+00:00</published><updated>2023-10-17T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/SSD</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/SSD/"><![CDATA[<h1 id="abstract">Abstract</h1>
<p>본 연구는 single deep neural network로 object detection을 수행한다. SSD라는 접근 방식은 bounding box를 피처맵의 위치별로 종횡비와 배율이 다른 default box 집합으로 이산화하는 것이다. 네트워크는 default box내에 있는 각각의 카테고리에
대해 점수를 부여하고 객체의 형태에 맞도록 box를 조정한다. 또한, 다양한 resolution을 갖는 피처맵을 결합하여 object의 사이즈가 다양하더라도 예측을 잘 수행할 수 있게한다. SSD는 region proposals이나 features resampling 단계를 없애고 모든 계산을
단일 네트워크에서 캡슐화 하기때문에 object proposals이 필요한 다른 네트워크에 비해 간단하다. PASCAL VOC, COCO, ILSVRC datsets에서 실험한 결과, SSD는 region proposals이 필요한(RCNN 등) 다른 모델들에 비해 준수한 정확도를 가진다. VOC 2007
에서 300 x 300 이미지를 사용했을 때 74.3% mAP, 59FPS를 기록했고, 512 x 512 이미지에서는 76.9% mAP를 기록했다. 대회 SOTA인 Fast R-CNN의 성능보다 더 좋은 성능이다. 또한, 다른 single stage methods와 비교해도 적은 input size로 더 좋은 정확도를 가졌다.</p>

<p><br /><br /></p>

<h1 id="details">Details</h1>
<h2 id="introduction">Introduction</h2>
<ul>
  <li>최근 object detection 모델은 bounding box를 제안 받고, 각 bbox에 대해 pixel or features를 resampling하고, 좋은 분류기를 연결하는 등의 접근 방식을 조금 변형해서 속도와 성능을 높였다. Fast R-CNN은 대표적인 모델로써 PASCAL, COCO 등의 탐지 대회에서 좋은 성능을 가졌다.</li>
  <li>이러한 접근 방식은 정확하긴하지만 계산 비용이 너무 크고 속도가 너무 느리다는 단점을 가지고 있다. 속도는 프레임당 초를 의미하는 SFP로 비교하는데 가장 빠르다는 Fast R-CNN조차 7 FPS를 가지고 있다. 현재는 속도를 얻는 대가로 정확도를 희생시키는게 전부인 추세다.</li>
  <li>본 논문은 bounding box 가설을 위한 pixel or featrues resampling을 수행하지 않아 속도를 크게 향상 시킨 deep neural network 기반의 object detection을 소개한다. 속도 향상을 위해 위의 작업을 제거한 것이 처음은 아니지만 몇가지 개선사항을 추가해 이전보다 정확도를 크게 높일 수 있었다.</li>
  <li>개선 사항으로는 작은 convolution filter를 사용하여 bounding box 위치의 offset과 object categories를 예측하는데 이 filter는 분리되어서 각각 다른 종횡비로 탐지를 수행하고 분리된 filter를 통해 여러개의 피처맵을 결합하여 다양한 scale의 object도 잘 탐지할 수 있도록했다.</li>
  <li>특히, 각각 다른 scale 정보를 가지고 예측을 수행하도록 multiple layer를 만들면 상대적으로 낮은 resolution input을 가지고도 높은 정확도와 빠른 속도를 낼 수 있었다.</li>
  <li>SSD 연구가 기여할 수 있는 부분은 다음과 같다.
    <ul>
      <li>Single shot detetcor의 이전 연구이자 SOTA였던 YOLO보다 빠르고, region proposals과 pooling을 사용한 다른 느린 모델(Fast R-CNN 포함)보다 더 정확한 single shot detector(SSD) for multiple categories를 소개한다.</li>
      <li>SSD의 핵심은 피처맵에 적용된 작은 convolutional filter를 사용하는 fixed default box 집합을 통해 category scores와 bbox의 offsets을 예측하는 것이다.</li>
      <li>정확도를 높이기 위해 다양한 scale의 피처맵에서 다양한 scale의 예측을 수행하고 종횡비별로 분리된 예측을 수행했다.</li>
      <li>이러한 구조로 단일 네트워크를 구축하고 low resolution input에서도 높은 정확도를 낼 수 있었고 속도와 정확도의 trade-off도 개선시켰다.</li>
      <li>실험 파트에서는 속도와 정확도를 평가지표로 삼아 PASCAL, COCO, ILSVRC datasets에서 각각 다른 image size를 사용했을 때 결과가 어떻게 달라지는지, 다른 SOTA 모델들과 어떤 차이가 있는지 실험했다.</li>
    </ul>
  </li>
</ul>

<p><br /><br /></p>

<h2 id="the-single-shot-detectorssd">The Single SHot Detector(SSD)</h2>
<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/591652b0-afd2-4b4c-8322-ac62acd483fb" />
  </p>
</div>

<p><br /></p>

<p>위의 그림은 SSD의 프레임워크를 나타낸 것으로, SSD는 CNN 단계에서 각각의 다른 크기의 피처맵을 통해 위치마다 다른 종횡비를 가진 몇개의 default box들을 평가한다. 또한, default box의 예측값으로 default box의 shape offsets과 confidences를 출력한다.
예를 들어, 먼저 default box들과 정답 box를 비교할 때 Figure 1.을 보면 8x8 피처맵에서는 2개의 default box가 고양이를, 4x4 피처맵에서는 하나의 default box가 개를 예측했고, 8x8 피처맵에서는 상대적으로 크기가 큰 개에 대해서는 예측을 잘 수행하지 못했다. 이렇게 default box 중 물체가 존재한다고 판단한 box를
positive, 그 외의 박스는 negative라고 할 수 있다. 모델의 loss는 localization loss(confidence loss(ex. softmax) 포함)를 가중합하여 계산한다.</p>

<p>추가적으로, 피처맵의 크기가 작다는 것은 그만큼 큰 filter를 썼다는 것이기때문에 크기가 작은 피처맵에서는 큰 객체에 대한 탐지가 가능하다. SSD에서는 이처럼 피처맵의 크기를 layer마다 다르게해서 다양한 scale의 object를 잘 탐지하도록 한다.</p>

<h3 id="1-model">1. Model</h3>
<p>SSD는 CNN을 기반으로 고정된 크기의 bbox들을 처리하고, bbox 내 존재하는 object들의 존재확신도 점수를 매긴다. 마지막 detection에서는 NMS 기법을 통해 중복된 박스를 제거하여 객체당 하나의 박스만 존재하도록 한다. 네트워크 앞단에는 image classification에서 높은 성능을 자랑했던 high quailty 구조를 사용(base network)하고 그 후 layer들을 추가하여 detection tasks를 수행하도록 한다.</p>

<p><br /></p>

<p><strong>Multi-scale feature maps for detection</strong></p>
<ul>
  <li>Base network에 CNN 구조를 연결한다. 새로 연결한 layers 에서는 피처맵의 사이즈를 점점 줄여나가면서 multiple scales에서 예측을 수행하도록 한다.</li>
  <li>각각의 feature layer마다 다른 scale의 예측을 수행하게 된다.</li>
</ul>

<p><br /></p>

<p><strong>Convolutional predictors for detection</strong></p>
<ul>
  <li>추가된 layer(아니면 base network에서 일부 layer를 더 활용해도됨)에서는 filter를 통해 각 층마다 고정된 detection predictions 집합(Figure 2.를 보면 conv 층 3x3 필터에 고정된 갯수의 default box가 생성된 것을 알 수 있음)을 생성할 수 있다.</li>
  <li>Layer에서는 mxnxp의 사이즈를 갖는 피처맵이 도출되는데, 기본적으로는 detection을 위해 3x3xp filter를 통해 cagteory scores와 default box coordinates를 계산한다.</li>
  <li>Bounding box의 offsets은 각 피처맵 위치에 존재하는 default box position과의 상대적 거리를 통해 계산된다(YOLO에서는 offsets을 구하기위해 CNN이 아닌 FC layer를 사용했음).</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/446069e6-a3fe-4428-90c1-fd92d90ff52d" />
  </p>
</div>

<p><br /></p>

<p><strong>Defalut boxes and aspect ratios</strong></p>
<ul>
  <li>Default box들의 집합은 네트워크 상단에 있는 여러개의 피처맵 셀들에 연결된다. 이 box들은 피처맵 셀의 알맞은 위치에 고정된다. 그리고 각 피처맵 셀에서 default box의 모양을 기준으로 offsets을 예측하고, 각 box 내 object들의 물체 존재 확신도를 예측한다.</li>
  <li>한 픽셀마다 k개 box들의 offsets 위치 4개와 c개의 클래스의 물체 존재 확신도 점수를 계산한다. 그 결과, mxn 피처맵에 대해 (c + 4)kmn의 결과가 도출된다.</li>
  <li>Default box는 Faster R-CNN의 anchor boxes와 비슷한 개념이지만 다양한 resolution을 가진 피처맵을 사용했다는 점에서 차이가 존재한다.</li>
</ul>

<p><br /></p>

<h3 id="2-training">2. Training</h3>
<p>SSD와 다른 detector와의 region proposals 단계에서 가장 큰 차이점은 SSD는 ground truth 정보가 고정된 detector outputs 집합에 꼭 할당 되어야한다는 것이다. 이 방법은 YOLO, Faster R-CNN, MultiBox 등 다른 모델들에서도 사용된다.
이 정보가 한번 할당되면 loss function과 역전파 전 과정에 적용된다. 학습에는 default boxes의 갯수와 scales을 얼마나 정할 것인지와 hard negative mining, data augmentation 등이 포함된다.</p>

<p><br /></p>

<p><strong>Matching strategy</strong></p>
<ul>
  <li>학습을 할 때는 어떤 default box가 ground truth와 관련이 있는지 파악하고, 이에 따라 네트워크를 훈련해야한다. 이는 ground truth box와 가장 많이 겹치는 default box를 선택한다. Multibox 모델과는 달리, 어떤 ground truth냐와는 상관없이 jaccard overlap(IOU)이 threshold인 0.5보다 큰 default box라면 모두 정답과 연관이 있다고 판단한다.</li>
  <li>이 방법은 가장 overlap이 큰 box를 하나만 남겨두는 것보다 오히려 다양한 scale을 가진 default box들을 통해 예측을 수행하여 learning problem을 단순화시킨다.</li>
</ul>

<p><br /></p>

<p><strong>Training objective</strong></p>
<ul>
  <li>SSD의 손실 함수는 MultiBox 모델과 유사하지만 multiple object categories를 처리하도록 수정되었다.</li>
  <li>$x_{ij}^p = [1,0]$를 p라는 category에 대해 i번째 default box와 j번째 ground truth가 연결되었는지를 나타낸다고 하자.</li>
  <li>default box가 여러개 있기때문에 ground box와 일치된다고 판단된 경우가 많다면 $\Sigma_{i} x_{ij}^p \geq 1$ 이 될 수 있다.</li>
  <li>전체 손실함수는 다음과 같이 정의된다.</li>
</ul>

\[L(x,c,l,g) = \frac{1}{N} (L_{conf}(x,c) + \alpha L_{loc}(x,l,g))\]

<ul>
  <li>N은 matching된 default box의 갯수를 의미하고 N = 0이면, loss를 0으로 설정한다.</li>
  <li>Localization loss는 predicted box($l$)과 ground truth box($g$)간의 Smooth L1 loss를 적용한다.</li>
  <li>Faster-RCNN과 비슷하게 default box($d$)에 대해 center($cx, cy$)와 $w,h$를 offsets으로 사용한다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/2de4931f-5d53-426b-9e70-57424ef50681" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Confidence loss는 multiple class confidences($c$)에 대한 softmax를 활용한다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/a8f7a2b1-9fa8-445a-92ae-af15c268d30c" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Cross validation을 통해 가중치항 $\alpha$는 1로 설정한다.</li>
  <li>정리하자면 손실함수는 ground truth와 매칭되는 default box들의 손실로 계산되는데
    <ul>
      <li>localization에 대한 손실을 offsets으로 계산하고</li>
      <li>box내 예측한 여러 class에 대한 손실을 softmax로 계산하여 두 손실을 합치는 방법을 사용한다.</li>
      <li>합쳐진 두 손실은 matching된 box의 수로 나눠 최종 loss를 구한다.</li>
    </ul>
  </li>
</ul>

<p><br /></p>

<p><strong>Choosing scales and aspect ratios for default boxes</strong></p>
<ul>
  <li>Object의 다양한 scale을 처리하기위해 어떤 모델은 이미지 사이즈를 직접적으로 처리하는 방법을 사용하지만 SSD에서는 하나의 단일 네트워크에서 layer마다 다른 scale을 갖는 피처맵을 활용하고 이를 공유한다.</li>
  <li>이전 연구들에 따라 층이 얕을수록 더 좁고 세밀한 정보를 파악할 수 있기때문에 얕은층과 깊은층의 피처맵들을 detection에 사용했다. 위의 Figure 1.에서 8x8과 4x4 피처맵이 그 예시이다.</li>
  <li>층마다 피처맵의 크기가 다르다면 층마다 receptive field가 달라지게될텐데 다행히 SSD의 default box는 꼭 실제 receptive fields와 box의 크기를 일치시킬 필요는 없다. 피처맵이 특정 objects의 scale에 맞게 학습되도록 default box의 크기를 조정할 수 있기때문이다. 특정 예측을 위해 m개의 피처맵이 필요하고 각 층을 k라고 한다면 default box의 scale을 다음과 같이 계산할 수 있다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/744ddda2-6713-4fbc-ba5d-b46d304cd090" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>$S_{min}$이 0.2, $S_{max}$가 0.9라는 것은 가장 얕은 층의 scale이 0.2, 가장 깊은 층의 scale이 0.9라는 의미이다. 모든 층의 scale은 이 안에 있다.</li>
  <li>다양한 default box의 aspect ratio는 $a_r \in [1,2,3,1/2,1/3]$ 으로 정의하고(1은 가로세로 비율이 같은 것, 2는 세로가 가로보다 2배 큰 것, 1/2는 가로가 세로보다 2배 큰 것) box마다 width는 $w_a^k = s_k \sqrt{a_r}$, height는 $h_a^k = s_k \sqrt{a_r}$로 계산한다.</li>
  <li>이러한 방법으로 default box의 크기를 유연하게 조절하며 다양한 scale에서 bounding box를 그리고 추출된 수많은 피처맵들을 통합하기때문에 특정 피처맵에서 객체를 발견하지 못하더라도 다른 피처맵에서 객체를 발견하여 예측에 큰 도움을 준다.</li>
</ul>

<p><br /></p>

<p><strong>Hard negative mining</strong></p>
<ul>
  <li>Matching 단계 이후에는 전경보다 배경이 훨씬 많기때문에 많은 default box들이 negative일 것이다.</li>
  <li>모든 negative examples을 쓰는 대신 각 default box마다 confidence loss를 정렬하고 가장 좋은 box를 선정해서 negative와 positive의 비율이 최대 3:1 정도가 되도록 한다.</li>
  <li>이를 통해 최적화가 더 빨라지고 안정적인 학습이 이루어질 수 있다.</li>
</ul>

<p><br /></p>

<p><strong>Data augmentation</strong></p>
<ul>
  <li>모델을 더 robust하게 만들기 위해서 input object의 크기나 형태를 변형하고 다음과 같은 option으로 input image를 랜덤하게 샘플링했다.
    <ul>
      <li>Original input image</li>
      <li>IOU가 0.1, 0.3, 0.5, 0.7, 0.9인 객체들을 패치로 샘플링</li>
      <li>랜덤하게 패치로 샘플링</li>
    </ul>
  </li>
  <li>patch의 크기는 원본 이미지의 0.1과 1사이로 하고 종횡비는 0.5와 2사이로 한다. 샘플링된 패치에도 ground truth box가 겹쳐있다면 예측에 사용한다.</li>
  <li>샘플링 단계 이후에는 샘플된 패치들을 고정된 크기로 리사이징하고 0.5의 확률로 수평변환시킨다. 추가로 photo-metric distortions을 적용한다.</li>
</ul>

<p><br /><br /></p>

<h2 id="experimental-results">Experimental Results</h2>
<p><strong>Base network</strong></p>
<ul>
  <li>사전훈련 모델로 VGG16을 사용하고 DeepLab-LargeFOV처럼 fc6와 fc7 layer를 convolutional layers로 변환시켰다. 또한, Pool5의 2x2(s2)를 3x3(s1)으로 바꾸고 atrous algorithm을 사용했다. 모든 dropout layers와 fc8 층을 제거했다.</li>
  <li>Fine tuning에는 optimizer SGD, learning rate 0.001, momentum 0.9, weight decay 0.0005, batch size 32를 사용했다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/446069e6-a3fe-4428-90c1-fd92d90ff52d" />
  </p>
</div>

<p><br /></p>

<h3 id="1-pascal-voc-2007">1. PASCAL VOC 2007</h3>
<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/6f2abeb2-41d2-45a9-98ba-e4209b023ea1" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Fast-RCNN, Faster R-CNN과 성능을 비교했다.</li>
  <li>SSD300은 이미지 사이즈를 300x300으로 조정한 모델이고 conv4_3, conv7(fc7), conv8_2, conv9_2, conv10_2, conv11_2에서 location과 confidences를 계산했다. Conv4_3의 default box scale은 0.1로 설정했고, 새로 추가된 convolutional layer의 파라미터는 xavier 방법으로 초기화했다.</li>
  <li>Conv4_3, conv10_2, conv11_2에는 각각의 피처맵 셀마다 4개의 default box를 사용했고, 다른 모든 layers에는 6개의 default box를 사용했다. 특히, conv4_3은 다른 layers들과 다른 feature scale을 가지고 있어서 L2 정규화로 feature norm을 20으로 조정하고 역전파 과정에서 scale을 학습했다.</li>
  <li>위의 Table 1을 보면 SSD300 모델이 Fast R-CNN보다 높은 정확도를 가진다는 것을 알 수 있고, SSD512로 이미지를 키웠을 때는 Faster R-CNN보다 높은 정확도를 기록했다. COCO 데이터를 추가하는 등 더 많은 data로 학습을 진행했을 때는 SSD512 모델에서 81.6% mAP를 기록했다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/12633d38-0345-4068-b1b8-10399f47e50a" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Figure 3.을 보면 white area(correct)가 넓은 것을 보아 SSD model이 다양한 category에 대해 예측을 잘 수행한 것을 알 수 있다. Recall도 85~90% 정도였고 IOU threshold를 0.1로 낮추면 더 높은 Recall을 보인다.</li>
  <li>SSD는 R-CNN처럼 두 가지의 step을 걸쳐 localization과 classification이 다르게 이루어지지 않고 하나의 네트워크에서 직접 bbox와 class를 예측하기때문에 R-CNN보다 localization error가 작다.</li>
  <li>하지만, 다양한 category들의 정보를 공유하기때문에 비슷한 objects에 대해서는 예측을 잘 수행하지 못한다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/d58bcee4-e629-4c27-876e-48b31017c8f4" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>Figure 4.를 보면 SSD는 bounding box size에 굉장히 민감한 것을 알 수 있고 큰 box보다 작은 box 때문에 성능이 더 하락했다. 작은 객체는 깊은 layer에서 정보를 얻어내기가 힘들기때문에 어찌보면 당연한 결과다(top layers에서는 세밀한 정보보다 주로 전체적인 context나 패턴을 파악하기때문에).</li>
  <li>300에서 512로 사이즈를 키워도 small object에 대한 탐지가 어려웠지만 그럼에도 다양한 scale의 박스와 피처맵을 통해 강건한 모델을 만들었다는 점에서 긍정적인 실험이었다.</li>
</ul>

<p><br /></p>

<h3 id="2-model-analysis">2. Model analysis</h3>
<p>SSD를 더 잘 이해하기위해 각각의 component들이 성능에 어떻게 영향을 미치는지 살펴보자. components의 영향을 잘 파악하기위해 모든 실험에 300 x 300사이즈의 이미지를 사용했다.</p>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/9d2d77d4-c332-4745-8317-fde689927ba5" />
  </p>
</div>

<p><br /></p>

<p><strong>Data augmentation is crucial</strong></p>
<ul>
  <li>Fast, Faster R-CNN은 원본이미지와 수평변환 이미지를 학습에 사용했다. SSD는 YOLO와 비슷한 방법으로 sampling을 수행해 성능을 향상시켰다.</li>
  <li>SSD의 샘플링 기법이 Fast, Faster R-CNN에도 효율적일진 모르겠지만 object의 robust에 중요한 단계인 classification 중간에 pooling이 있기때문에 아마 효과가 좋진않을 것이다.</li>
</ul>

<p><br /></p>

<p><strong>More default box shapes is better</strong></p>
<ul>
  <li>SSD는 피처맵 location마다 보통 6개의 default box를 사용했다. 만약 aspect ratios $\frac{1}{3}, 3$으로 boxes를 제거하면 성능이 하락했고, $\frac{1}{2}, 2$도 마찬가지였다.</li>
  <li>다양한 크기의 default boxes를 사용하는 것이 성능 향상에 도움이 된다.</li>
</ul>

<p><br /></p>

<p><strong>Atrous is faster</strong></p>
<ul>
  <li>DeepLab 모델처럼 VGG16에 atrous convolution을 적용했다.</li>
  <li>만약, VGG16을 그대로 사용해서 pool5(2x2, s2) 사용, fc6, fc7의 파라미터에 subsampling 적용 X, conv5_3을 추가하는 방법을 사용하면 정확도는 같은데 속도는 20% 느려졌다.</li>
</ul>

<p><br /></p>

<p><strong>Multiple output layers at different resolutions is better</strong>
<br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/a8746e22-5b66-413b-88f6-1527259cc8ed" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>SSD는 다양한 output layers에서 다양한 scale의 default box를 사용하는 것이 핵심이기때문에 layers를 점점 줄여가면서 이 효과를 검증해봤다.</li>
  <li>정확한 비교를 위해 layer를 제거할 때마다 default box tiling(각 픽셀마다 default box를 잘 배열하는 것)을 조정하여 총 상자 수를 원본과 비슷한 8732개로 유지했다(하지만, 모든 실험에서 tiling을 한 것은 아님). 남은 layer에 box scale을 더 많이 쌓고 필요한 경우 box scale을 조정했다. layer에 box를 쌓을 때 box가 이미지 경계에 있는 경우가 많으므로 주의하며 쌓아야한다.</li>
  <li>Table 3.는 layer가 적을수록 성능이 떨어지는 것을 보여준다.</li>
  <li>Faster R-CNN처럼 경계에 있는 box를 무시한채 예측을 진행했더니 흥미로운 사실을 발견했다. 만약 11_2, 10_2와 같은 깊은 layer를 사용한다면 성능이 크게 떨어진다는 것이다. 아마 이미지 경계에 있는 box들이 예측에 포함되지않아 큰 객체를 감쌀 수 있는 큰 box가 충분치 않아서일 것이다.</li>
  <li>또한, conv7만을 사용했을 때 최악의 성능이 나왔는데 이것으로 다양한 scale의 box를 여러 layer에 적절히 분배하는 것이 성능에 큰 영향을 끼친다는 것을 알 수 있다.</li>
  <li>SSD는 lower resolution(300x300) input image를 사용했는데도 Faster R-CNN과 비슷한 정확도를 가졌다.</li>
</ul>

<p><br /></p>

<h3 id="3-pascal-voc2012">3. PASCAL VOC2012</h3>
<ul>
  <li>VOC2007과 똑같은 setting으로 실험을 진행한 결과 같은 performance를 보였다. SSD 300은 Fast/Faster R-CNN보다 정확도가 높았고 SSD512는 Faster R-CNN보다 4.5% 더 높은 mAP를 기록했다.</li>
  <li>YOLO와 비교했을 때도 훨씬 좋은 성능을 가졌고, COCO dataset으로 추가학습을 하면 mAP가 80%였다.</li>
</ul>

<p><br /></p>

<h3 id="4-coco--ilsvrc">4. COCO &amp; ILSVRC</h3>
<ul>
  <li>COCO는 PASCAL보다 객체가 좀 더 작기때문에 모든 레이어에 small default box를 사용했다. 따라서 최소 scale 0.2를 0.15로 조정하고 conv4_3의 scale을 0.07로 조정했다.</li>
  <li>SSD512는 conv12-2를 추가했고, $S_{min}$을 0.1, conv4_3을 0.04의 scale로 조정했다.</li>
  <li>PASCAL과 비슷하게 SSD300이 Fast R-CNN보다 성능이 좋았지만 mAP@0.75에서는 Faster R-CNN과 성능이 비슷했고 mAP@0.5에서는 성능이 좀 더 낮았다.</li>
  <li>SSD512에서는 mAP@0.75와 mAP@0.5 모두에서 Faster R-CNN보다 성능이 좋았다. 하지만, 작은 객체에 대해서는 Faster R-CNN과 큰 격차를 내진 못했는데 아마 Faster R-CNN은 RPN(Region Proposals Network)과 Fast R-CNN(분류 및 bbox 예측) 두 단계에 걸쳐 box에 대한 세밀한 조정이 이루어지기때문에 작은 객체에 대해 잘 예측할 수 있었을 것으로 추측된다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/5497169b-8fd9-4a8f-b7d4-4f6f3514e2de" />
  </p>
</div>

<p><br /></p>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/ccaa82e4-c348-4bb5-9764-d5e3f8d8218f" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>ILSVRC에서도 COCO와 동일한 네트워크를 사용했고 그 결과 val2 set에서 43.4% mAP를 달성했다.</li>
</ul>

<p><br /></p>

<h3 id="5-data-augmentation-for-small-object-accuracy">5. Data Augmentation for Small Object Accuracy</h3>
<ul>
  <li>데이터 증강은 PASCAL VOC와 같은 작은 datasets에서 성능을 효과적으로 높일 수 있다. 이미지를 랜덤하게 자르는 방법이 ‘zoom in’의 효과처럼 크기가 큰 데이터들을 만들어낼 수 있는 것이다.</li>
  <li>크기가 작은 데이터를 만들어내기 위해 ‘zoom out’ 효과를 적용하려면 원본 이미지를 랜덤하게 자르기 전에 평균값으로 채워졌고 원본 이미지보다 16배 큰 캔버스에 이미지들을 랜덤하게 위치시킨다. 이렇게 zoom out 된 상태에서 crop을 진행하면 크기가 작은 data를 생성할 수 있는 것이다. 그 결과 여러 데이터셋에서 mAP가 2~3% 정도 증가했다.</li>
  <li>SSD를 개선하는 또 다른 방법은 default box의 tiling을 더 잘 디자인해서 receptive field와 default box가 잘 맞도록 하는 것이다. 이 task는 향후 연구를 위해 남겨둔다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/06aebf10-f169-416e-a89e-8cd2a8ecd428" />
  </p>
</div>

<p><br /></p>

<h3 id="6-inference-time">6. Inference time</h3>
<p><br /></p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/dc964ec2-6644-46a6-baaa-b74e8b6df83c" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>SSD에서는 수많은 box를 처리해야하기때문에 inference 과정에서 NMS를 잘 활용하는 것이 매우 중요하다. 먼저, confidence threshold를 0.01로 설정하면 많은 box들을 제거할 수 있고, 그 후 IOU threshold 0.45를 사용한 NMS를 통해 이미지당 200개의 box만 남도록 유지할 수 있다.</li>
  <li>NMS를 통해 SSD300은 VOC 200 classes에 대해 이미지당 1.7 msec 밖에 소요되지 않는다.</li>
  <li>Table.7을 보면 SSD300과 SSD512 모두 Faster R-CNN의 정확도와 속도를 능가한다. SSD300은 70% 이상의 mAP를 달성한 최초의 real-time method 이다.</li>
  <li>SSD 모델의 forward time 중 80%는 base network(VGG16)에서 소요된 것이기때문에 base network를 더 빠른 네트워크로 구축하면 속도면에서 향상되어 SSD512도 real-time으로 만들 수 있을 것이다.</li>
</ul>

<p><br /><br /></p>

<h2 id="related-work">Related Work</h2>
<ul>
  <li>이미지에서 객체를 탐지하는 방법은 sliding window를 활용하는 방법과 region proposals를 활용하는 방법이 있다. R-CNN이 성능을 크게 향상시키자 region proposals을 기반으로 한 object detection이 트렌드가 되었다.</li>
  <li>R-CNN은 다양한 방법을 통해 모델을 개선시켜나갔는데 첫번째 방법으로는 classification에 소요되는 시간을 줄이는 접근이다. SPPPnet에서는 spatial pyramid pooling layer를 적용하여 region size와 scale을 robust하게 만들었고, 여러 이미지 resolution을 통해 생성된 피처맵들을 classification layer에서 다시 사용하므로 속도 문제를 개선시켰다. Fast R-CNN은 SPPnet을 보완시켜서 모든 layers를 fine-tuning 시켜서 confidences와 bounding box regression loss를 최소화시켰다.</li>
  <li>R-CNN을 개선시킨 두번째 방법으로는 deep neural network를 사용하여 proposal quaility를 개선시키는 접근이다. MultiBox와 같은 최근 모델에서는 low-level image features를 기반으로 한 selective search region proposals을 deep neural network에서 직접적으로 생성되는 proposals로 대체했다. 이 방법은 정확도를 증가시켰지만 proprosals과 classification을 위해 종속된 두 가지 neural network를 학습시켜야했으므로 복잡하다는 단점이 있다.</li>
  <li>Faster R-CNN에서는 selective search를 RPN으로 대체하고 기존 Fast R-CNN에 RPN을 통합하여 CNN과 prediction layers가 정보를 잘 공유할 수 있도록 fine-tuning을 진행했다.</li>
  <li>SSD는 RPN과 비슷한 방법을 사용하는데 Faster R-CNN에서 anchor boxes를 사용한 것처럼 SSD는 고정된 default box 집합을 사용한다는 것이다. 하지만, Faster R-CNN에서는 pool feature를 추출하고 평가는 다른 분류기에서 이루어졌는데 SSD는 각각의 box마다 여러 object category에 대한 score를 계산할 수 있어 특징 추출과 평가가 동시에 이루어진다는 점에서 차이가 있다. 따라서, SSD는 여러 tasks들을 통합하여 Faster R-CNN보다 학습이 쉽고 빠른 모델을 만들 수 있었다.</li>
  <li>SSD와 비슷한 또 다른 접근 방식은 proposals 단계를 생략하고 여러 categories에 대해 bbox와 confidence를 바로 예측하는 것이다. Sliding window 기법의 심층 버전인 OverFeat은 
object categories의 confidence를 얻은 후에 각 위치마다 가장 좋은 피처맵을 사용해서 bounding box를 예측한다. YOLO는 가장 좋은 피처맵을 통해 multiple categories의 confidences와 bounding box를 예측한다.</li>
  <li>SSD는 이 두가지 모델처럼 proposals step이 없고 default box를 사용하여 confidence와 bbox를 예측하기때문에 비슷한 접근방식이라고 할 수 있다. 하지만, 다양한 scales의 피처맵 location에서 다양한 종횡비를 가진 bbox를 사용하기때문에 기존 접근 방식보다 더 유연한 방법이라고 할 수 있다.</li>
  <li>만약 default box를 가장 좋은 하나의 피처맵 location에서만 사용한다면 SSD는 OverFeat과 비슷한 구조를 가졌을 것이고, 가장 좋은 피처맵을 사용하면서 convoutional layer대신 fully connected layer를 추가하고 다양한 aspect ratios를 고려하지 않는다면 YOLO와 비슷한 모델이 될 것이다.</li>
</ul>

<p><br /><br /></p>

<h1 id="conclusion">Conclusion</h1>
<ul>
  <li>SSD는 multiple feature maps에 적용된 다양한 scale의 convolutional bounding box를 사용하는 fast single-shot object detector이다.</li>
  <li>많은 default box들이 성능을 향상시켰고 기존보다 몇배 더 많은 box predictions sampling location, scale, aspect ratio가 사용됐다.</li>
  <li>SSD512 모델은 기존 SOTA였던 Faster R-CNN보다 3배나 더 빠르면서 성능은 뛰어난 모델이고 SSD300은 58FPS로 YOLO보다 빠르고 정확한 real time 모델이다.</li>
  <li>SSD는 비디오에서 객체를 탐지하고 추적하는 모델의 일부로 사용되어 향후 연구에 좋은 영향을 끼칠 것이다.</li>
</ul>

<p><br /><br /></p>

<h1 id="개인적인-생각">개인적인 생각</h1>
<ul>
  <li>SSD는 단일 네트워크로 빠른 이미지 처리 뿐만 아니라 다양한 scale의 default box를 사용함으로써 다양한 객체의 크기에 유연하게 반응할 수 있다는 점에서 매우 의미있는 연구였다.</li>
  <li>YOLO와 비슷하게 SOTA 모델인 R-CNN 계열의 모델들이나 비슷하거나 다른 접근방식을 가진 다양한 모델들과 성능을 비교했고, 좋았던 부분이나 부족한 부분을 시사할 수 있었기때문에 이를 바탕으로 부족한 부분을 개선하면 향후 연구에도 큰 도움이 될 것이다.</li>
  <li>VGG보다 더 좋은 성능을 내는 GoogLeNet, ResNet과 같은 모델을 기반으로 pre-trained 시켰다면 본 논문에서도 언급한바와 같이 더 빠르고 정확한 모델이 만들어지진 않았을까라는 생각이 들었다.</li>
  <li>Multiple-scale의 default box를 사용하더라도 결국 작은 객체에 대해서는 탐지가 쉽지않았는데 향후 연구에서는 작은 객체를 잘 탐지하도록 어떤 기술을 사용할지 궁금하다.</li>
  <li>SSD와 YOLO 모두 one stage detector로 region proposals과 분류 및 box prediction이 단일 네트워크에서 가능하게해서 R-CNN과 같은 two stage detector보다 속도를 더 빠르게했다. 2개의 분리된 단계를 하나의 단계로 압축시켰기때문에 어찌보면 속도를 가장 잘 줄일 수 있는 방법을 사용했다고 볼 수 있는데 앞으로는 어떠한 기술들로 정확도를 크게 희생시키지않으면서 속도를 더 빠르게 할지, 더 나아가 정확도도 크게 개선시키고 속도도 크게 줄일지 기대가 된다.</li>
</ul>

<p><br /><br /></p>

<h1 id="구현">구현</h1>
<p>SSDLite 모델을 사용해보자(<a href="https://github.com/pytorch/vision/blob/main/torchvision/models/detection/ssdlite.py">코드 참고</a>). 데이터셋 및 파일 경로 설정은 <a href="https://hyeonseung0103.github.io/detection/Fast_and_Faster_RCNN/">Fast &amp; Faster RCNN 포스팅 구현 파트 참고</a>.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">json</span>
<span class="kn">import</span> <span class="nn">cv2</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">os</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
<span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">from</span> <span class="nn">pycocotools.coco</span> <span class="kn">import</span> <span class="n">COCO</span>
<span class="kn">from</span> <span class="nn">PIL</span> <span class="kn">import</span> <span class="n">Image</span>
<span class="kn">import</span> <span class="nn">time</span>
<span class="kn">import</span> <span class="nn">transforms</span> <span class="k">as</span> <span class="n">T</span>
<span class="kn">from</span> <span class="nn">engine</span> <span class="kn">import</span> <span class="n">train_one_epoch</span><span class="p">,</span> <span class="n">evaluate</span>
<span class="kn">import</span> <span class="nn">utils</span>
<span class="kn">from</span> <span class="nn">drive.MyDrive.paper_practice.custom_dataset.soccer_dataset</span> <span class="kn">import</span> <span class="n">SoccerDataset</span> <span class="c1"># SoccerDataset 모듈화
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">albumentations</span> <span class="k">as</span> <span class="n">A</span>
<span class="kn">import</span> <span class="nn">transforms</span> <span class="k">as</span> <span class="n">T</span>

<span class="k">def</span> <span class="nf">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="p">):</span>
    <span class="n">transforms</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="c1"># if train:
</span>    <span class="c1">#     transforms.append(A.HorizontalFlip(0.5))
</span>    <span class="c1">#     transforms.append(A.VerticalFlip(0.5))
</span>    <span class="k">return</span> <span class="n">A</span><span class="p">.</span><span class="n">Compose</span><span class="p">(</span><span class="n">transforms</span><span class="p">,</span> <span class="n">bbox_params</span><span class="o">=</span><span class="n">A</span><span class="p">.</span><span class="n">BboxParams</span><span class="p">(</span><span class="nb">format</span><span class="o">=</span><span class="s">'pascal_voc'</span><span class="p">,</span> <span class="n">label_fields</span><span class="o">=</span><span class="p">[</span><span class="s">'labels'</span><span class="p">]))</span>
        <span class="c1"># label_fields는 호출할 때 입력한 key와 맞아야함. key이름을 labels로 했으니까 field로 labels
</span>        <span class="c1"># 이미 x2,y2형식으로 바꿨으니까 coco가 아닌 pascal 형식
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torchvision</span>
<span class="kn">from</span> <span class="nn">torchvision.models.detection.ssdlite</span> <span class="kn">import</span> <span class="n">SSDLiteClassificationHead</span>
<span class="kn">from</span> <span class="nn">torchvision.models.detection</span> <span class="kn">import</span> <span class="n">_utils</span> <span class="k">as</span> <span class="n">det_utils</span>
<span class="kn">from</span> <span class="nn">functools</span> <span class="kn">import</span> <span class="n">partial</span>
<span class="kn">from</span> <span class="nn">torch</span> <span class="kn">import</span> <span class="n">nn</span>

<span class="n">model</span> <span class="o">=</span> <span class="n">torchvision</span><span class="p">.</span><span class="n">models</span><span class="p">.</span><span class="n">detection</span><span class="p">.</span><span class="n">ssdlite320_mobilenet_v3_large</span><span class="p">(</span><span class="n">pretrained</span> <span class="o">=</span> <span class="bp">True</span><span class="p">)</span>
<span class="n">num_classes</span> <span class="o">=</span> <span class="mi">3</span> <span class="c1"># has ball, no ball, background
</span>
<span class="c1"># backbone 끝단의 output을 head의 input으로 사용하고 이미지 크기는 640x640으로
</span><span class="n">in_channels</span> <span class="o">=</span> <span class="n">det_utils</span><span class="p">.</span><span class="n">retrieve_out_channels</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">backbone</span><span class="p">,</span> <span class="p">(</span><span class="mi">640</span><span class="p">,</span> <span class="mi">640</span><span class="p">))</span>
<span class="n">num_anchors</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="n">anchor_generator</span><span class="p">.</span><span class="n">num_anchors_per_location</span><span class="p">()</span> <span class="c1"># 기존 앵커 그대로 사용
</span><span class="n">norm_layer</span> <span class="o">=</span> <span class="n">partial</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">,</span> <span class="n">eps</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span> <span class="n">momentum</span><span class="o">=</span><span class="mf">0.03</span><span class="p">)</span> <span class="c1"># 정규화 layer 설정
</span>
<span class="c1"># regression은 위치만 보기 때문에 class 정보가 필요하지 않음. 앵커도 기존의 앵커 비율을 사용했기때문에 굳이 수정 X
</span><span class="n">model</span><span class="p">.</span><span class="n">head</span><span class="p">.</span><span class="n">classification_head</span> <span class="o">=</span> <span class="n">SSDLiteClassificationHead</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">num_anchors</span><span class="p">,</span> <span class="n">num_classes</span><span class="p">,</span> <span class="n">norm_layer</span><span class="p">)</span>

<span class="n">device</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span> <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">cuda</span><span class="p">.</span><span class="n">is_available</span><span class="p">()</span> <span class="k">else</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cpu'</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># forward 테스트
</span><span class="n">a</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">,</span> <span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
<span class="n">dl</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
  <span class="n">collate_fn</span><span class="o">=</span><span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>

<span class="n">images</span><span class="p">,</span><span class="n">targets</span> <span class="o">=</span> <span class="nb">next</span><span class="p">(</span><span class="nb">iter</span><span class="p">(</span><span class="n">dl</span><span class="p">))</span>
<span class="n">images</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">image</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">)</span>
<span class="n">targets</span> <span class="o">=</span> <span class="p">[{</span><span class="n">k</span><span class="p">:</span> <span class="n">v</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">t</span><span class="p">.</span><span class="n">items</span><span class="p">()}</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">targets</span><span class="p">]</span>

<span class="n">output</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span><span class="n">targets</span><span class="p">)</span>
<span class="n">output</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 데이터 정의
</span><span class="n">train_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">,</span> <span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">))</span>
<span class="n">val_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">VAL_DATA_PATH</span><span class="p">,</span> <span class="n">VAL_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">train_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">train_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>

<span class="n">val_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">val_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="c1"># 재연성을 위해 셔플 False
</span>                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">num_epochs</span> <span class="o">=</span> <span class="mi">50</span>
<span class="n">val_loss_tmp</span> <span class="o">=</span> <span class="mi">10000</span>
<span class="n">best_epoch_tmp</span> <span class="o">=</span> <span class="mi">1</span>
<span class="n">early_stopping_cnt</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">early_stop</span> <span class="o">=</span> <span class="mi">20</span>

<span class="n">params</span> <span class="o">=</span> <span class="p">[</span><span class="n">p</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">parameters</span><span class="p">()</span> <span class="k">if</span> <span class="n">p</span><span class="p">.</span><span class="n">requires_grad</span><span class="p">]</span>
<span class="n">optimizer</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">SGD</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span>
                            <span class="n">momentum</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">weight_decay</span><span class="o">=</span><span class="mf">0.0005</span><span class="p">)</span>

<span class="c1"># lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer,
#                                                 step_size=3,
#                                                 gamma=0.9)
</span>
<span class="c1">#lr_scheduler = torch.optim.lr_scheduler.MultiplicativeLR(optimizer=optimizer, lr_lambda=lambda lr: 0.95 ** lr)
</span><span class="n">lr_scheduler</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">optim</span><span class="p">.</span><span class="n">lr_scheduler</span><span class="p">.</span><span class="n">CosineAnnealingLR</span><span class="p">(</span><span class="n">optimizer</span><span class="p">,</span> <span class="n">T_max</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">eta_min</span><span class="o">=</span><span class="mf">0.0001</span><span class="p">)</span>

<span class="k">print</span><span class="p">(</span><span class="s">'----------------------train start--------------------------'</span><span class="p">)</span>

<span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">num_epochs</span><span class="o">+</span><span class="mi">1</span><span class="p">):</span>
  <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
  <span class="n">model</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
  <span class="n">epoch_loss</span> <span class="o">=</span> <span class="mi">0</span>
  <span class="n">prog_bar</span> <span class="o">=</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">train_data_loader</span><span class="p">,</span> <span class="n">total</span><span class="o">=</span><span class="nb">len</span><span class="p">(</span><span class="n">train_data_loader</span><span class="p">))</span>

  <span class="k">for</span> <span class="n">images</span><span class="p">,</span> <span class="n">targets</span> <span class="ow">in</span> <span class="n">prog_bar</span><span class="p">:</span>
    <span class="n">images</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">image</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">)</span>
    <span class="n">targets</span> <span class="o">=</span> <span class="p">[{</span><span class="n">k</span><span class="p">:</span> <span class="n">v</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">t</span><span class="p">.</span><span class="n">items</span><span class="p">()}</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">targets</span><span class="p">]</span>

    <span class="n">loss_dict</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span> <span class="n">targets</span><span class="p">)</span>
    <span class="n">optimizer</span><span class="p">.</span><span class="n">zero_grad</span><span class="p">()</span> <span class="c1"># optimization 정보가 누적되지않도록 초기화
</span>    <span class="n">loss</span> <span class="o">=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">loss</span> <span class="k">for</span> <span class="n">loss</span> <span class="ow">in</span> <span class="n">loss_dict</span><span class="p">.</span><span class="n">values</span><span class="p">())</span>
    <span class="n">loss</span><span class="p">.</span><span class="n">backward</span><span class="p">()</span>
    <span class="n">optimizer</span><span class="p">.</span><span class="n">step</span><span class="p">()</span>
    <span class="n">epoch_loss</span> <span class="o">+=</span> <span class="n">loss</span><span class="p">.</span><span class="n">item</span><span class="p">()</span>
  <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'epoch : </span><span class="si">{</span><span class="n">epoch</span><span class="si">}</span><span class="s">, Loss : </span><span class="si">{</span><span class="n">epoch_loss</span><span class="si">}</span><span class="s">, time : </span><span class="si">{</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>

  <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">epoch_val_loss</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="n">val_start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
    <span class="k">for</span> <span class="n">images</span><span class="p">,</span> <span class="n">targets</span> <span class="ow">in</span> <span class="n">val_data_loader</span><span class="p">:</span>
        <span class="n">images</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">image</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">)</span>
        <span class="n">targets</span> <span class="o">=</span> <span class="p">[{</span><span class="n">k</span><span class="p">:</span> <span class="n">v</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">t</span><span class="p">.</span><span class="n">items</span><span class="p">()}</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">targets</span><span class="p">]</span>

        <span class="n">val_loss_dict</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">images</span><span class="p">,</span> <span class="n">targets</span><span class="p">)</span>
        <span class="n">epoch_val_loss</span> <span class="o">+=</span> <span class="nb">sum</span><span class="p">(</span><span class="n">loss</span> <span class="k">for</span> <span class="n">loss</span> <span class="ow">in</span> <span class="n">val_loss_dict</span><span class="p">.</span><span class="n">values</span><span class="p">())</span>

    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'Val Loss : </span><span class="si">{</span><span class="n">epoch_val_loss</span><span class="si">}</span><span class="s">, time : </span><span class="si">{</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">val_start</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>
    <span class="k">if</span> <span class="n">epoch_val_loss</span> <span class="o">&lt;</span> <span class="n">val_loss_tmp</span><span class="p">:</span> <span class="c1"># best 모델만 저장
</span>        <span class="n">early_stopping_cnt</span> <span class="o">=</span> <span class="mi">0</span>
        <span class="n">best_epoch_tmp</span> <span class="o">=</span> <span class="n">epoch</span>
        <span class="n">val_loss_tmp</span> <span class="o">=</span> <span class="n">epoch_val_loss</span>
        <span class="n">torch</span><span class="p">.</span><span class="n">save</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">state_dict</span><span class="p">(),</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">WEIGHTS_PATH</span><span class="si">}</span><span class="s">ssd_</span><span class="si">{</span><span class="n">num_epochs</span><span class="si">}</span><span class="s">.pt'</span><span class="p">)</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="n">early_stopping_cnt</span> <span class="o">+=</span> <span class="mi">1</span> <span class="c1"># 손실이 늘었으면 early_stopping count
</span>    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'현재까지 best 모델은 Epochs </span><span class="si">{</span><span class="n">best_epoch_tmp</span><span class="si">}</span><span class="s">번째 모델입니다.'</span><span class="p">)</span>

  <span class="k">if</span> <span class="n">early_stopping_cnt</span> <span class="o">==</span> <span class="n">early_stop</span><span class="p">:</span>
    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">early_stop</span><span class="si">}</span><span class="s">번 동안 validation 성능 개선이 없어 학습을 조기 종료합니다.'</span><span class="p">)</span>
    <span class="k">break</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># test 데이터 정의 및 평가
</span><span class="n">test_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span><span class="p">,</span> <span class="n">TEST_LAB_PATH</span><span class="p">,</span> <span class="n">get_transforms</span><span class="p">(</span><span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>

<span class="n">test_data_loader</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">utils</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">DataLoader</span><span class="p">(</span><span class="n">test_dataset</span><span class="p">,</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
                                                <span class="n">collate_fn</span> <span class="o">=</span> <span class="n">utils</span><span class="p">.</span><span class="n">collate_fn</span><span class="p">)</span>
<span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">test_data_loader</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 결과 시각화
</span><span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">test_dataset</span><span class="p">[</span><span class="mi">100</span><span class="p">]</span>
<span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
<span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">prediction</span> <span class="o">=</span> <span class="n">model</span><span class="p">([</span><span class="n">i</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)])[</span><span class="mi">0</span><span class="p">]</span>

<span class="n">i</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">i</span><span class="p">.</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>
<span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">]):</span>
  <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">cpu</span><span class="p">(),</span> <span class="n">dtype</span> <span class="o">=</span> <span class="nb">int</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
  <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'labels'</span><span class="p">][</span><span class="n">idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/958d5e63-09d6-4548-860f-5e4abd033c70" />
  </p>
  <p>nms 적용전 SSD 결과</p>
</div>

<p><br /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="nn">torchvision.ops</span> <span class="kn">import</span> <span class="n">nms</span>

<span class="n">selected_idx</span> <span class="o">=</span> <span class="n">nms</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">],</span> <span class="n">prediction</span><span class="p">[</span><span class="s">'scores'</span><span class="p">],</span> <span class="n">iou_threshold</span> <span class="o">=</span> <span class="mf">0.2</span><span class="p">)</span>
<span class="n">selected_boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'boxes'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>
<span class="n">selected_labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'labels'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>
<span class="n">selected_scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">prediction</span><span class="p">[</span><span class="s">'scores'</span><span class="p">])[</span><span class="n">selected_idx</span><span class="p">]</span>

<span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">test_dataset</span><span class="p">[</span><span class="mi">100</span><span class="p">]</span>
<span class="n">i</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">i</span><span class="p">.</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>
<span class="k">for</span> <span class="n">idx</span><span class="p">,</span><span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">selected_boxes</span><span class="p">):</span>
  <span class="k">if</span> <span class="n">selected_scores</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span> <span class="o">&gt;</span> <span class="mf">0.9</span><span class="p">:</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">cpu</span><span class="p">(),</span> <span class="n">dtype</span> <span class="o">=</span> <span class="nb">int</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">i</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">selected_labels</span><span class="p">[</span><span class="n">idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">i</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/d263f8c3-c844-487f-b206-2c9b87f32967" />
  </p>
  <p>nms 적용후 SSD 결과(confidence score가 0.9 이상일 때)</p>
</div>

<p><br /></p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/095168fd-e777-4331-9518-f1c780bb01a3" />
  </p>
  <p>Faster R-CNN 결과</p>
</div>

<p><br /></p>

<p>SSD의 가벼운 버전인 SSDLite를 사용했을 때 test set에서 mAP@0.5:0.95이 0.546이 나왔다. 이는 Faster R-CNN의 0.407보다 훨씬 높은 수치이다. 위의 이미지를 통해 SSD가 Faster R-CNN보다 localization과 classification을 모두 잘 수행한 것을 알 수 있다. SSD에서는 볼을 소유하고 있는 사람에게 1의 클래스를 잘 부여했지만, Faster R-CNN은 모든 사람을 볼을 소유하고 있지 않은 2의 클래스로 예측했다. 학습시간 또한 SSD는 한 에포크당 약 20초의 학습 시간이 걸렸는데 Faster R-CNN이 에포크당 3분 30초가 걸린 것과 비교하면 약 10배 이상 빠른 속도다.</p>

<p>논문에서 언급한대로 당시 SOTA model이었던 Faster R-CNN보다 정확도와 속도 모든 면에서 우월한 SSD였다.</p>

<h1 id="이미지-출처">이미지 출처</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1512.02325v5.pdf">SSD: Single Shot MultiBox Detector Paper</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[Abstract 본 연구는 single deep neural network로 object detection을 수행한다. SSD라는 접근 방식은 bounding box를 피처맵의 위치별로 종횡비와 배율이 다른 default box 집합으로 이산화하는 것이다. 네트워크는 default box내에 있는 각각의 카테고리에 대해 점수를 부여하고 객체의 형태에 맞도록 box를 조정한다. 또한, 다양한 resolution을 갖는 피처맵을 결합하여 object의 사이즈가 다양하더라도 예측을 잘 수행할 수 있게한다. SSD는 region proposals이나 features resampling 단계를 없애고 모든 계산을 단일 네트워크에서 캡슐화 하기때문에 object proposals이 필요한 다른 네트워크에 비해 간단하다. PASCAL VOC, COCO, ILSVRC datsets에서 실험한 결과, SSD는 region proposals이 필요한(RCNN 등) 다른 모델들에 비해 준수한 정확도를 가진다. VOC 2007 에서 300 x 300 이미지를 사용했을 때 74.3% mAP, 59FPS를 기록했고, 512 x 512 이미지에서는 76.9% mAP를 기록했다. 대회 SOTA인 Fast R-CNN의 성능보다 더 좋은 성능이다. 또한, 다른 single stage methods와 비교해도 적은 input size로 더 좋은 정확도를 가졌다.]]></summary></entry><entry><title type="html">YOLO : You Only Look Once: Unified, Real-Time Object Detection 논문 요약</title><link href="https://kimyeonz.github.io/blog/detection/YOLO/" rel="alternate" type="text/html" title="YOLO : You Only Look Once: Unified, Real-Time Object Detection 논문 요약" /><published>2023-10-16T00:00:00+00:00</published><updated>2023-10-16T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/YOLO</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/YOLO/"><![CDATA[<h1 id="abstract">Abstract</h1>
<p>객체 탐지의 이전 연구들에서는 분류기에서도 detection을 수행할 수 있도록 했다. 본 연구에서는 하나의 단일 신경망을 사용해서 여러 bounding box와 class probabilities를 예측하는 객체 탐지를 regression 문제로 취급했다. 본 연구의 YOLO 모델은 객체 탐지 파이프라인 전체가 하나의 네트워크로 이루어져있어서 실시간 이미지를 초당 45프레임으로 처리할만큼 굉장히 빠르다. 다른 네트워크와 비교했을 때 localization에 대한 오류가 존재하긴하지만 배경에 대해서는 예측을 잘 수행하고 다른 도메인에 대해서도 RCNN, DPM보다 일반화가 잘 된 예측을 수행한다.</p>

<p><br /><br /></p>

<h1 id="details">Details</h1>
<h2 id="introduction">Introduction</h2>
<ul>
  <li>최근 detection networks들은 분류기가 탐지를 수행할 수 있도록 변형한다. 분류기는 탐지 역할을 하기 위해서 다양한 위치와 크기의 이미지들에서 객체를 탐지하고 평가한다. DPM 형태의 모델에서는
sliding window 개념을 활용하여 전체 이미지에 대해 균일한 간격으로 분류를 수행한다.</li>
  <li>보다 최근인 RCNN에서는 물체가 있을법한 곳에 bounding box를 생성하고 분류기는 이 bounding box에 대해 분류를 수행한다. 분류 후에는 bounding box 내의 객체가 중복되면
이를 제거하고, box를 다른 object들을 기반으로 rescoring 한다. RCNN의 경우 각각의 tasks들이 별개로 이루어져있기때문에 복잡하고 예측에 많은 시간이 소요된다는 단점을 가지고있다.</li>
  <li>YOLO(You Only Look Once)에서는 단일 네트워크로 CNN에서 여러 bbox와 각 bbox 마다 존재하는 class의 확률을 예측한다. 아래 Figure 1. 이미지를 보면, YOLO는 이미지를 고정된 448 x 448
사이즈로 리사이징 하고, CNN을 통해 class 분류 및 bbox를 생성하고, confidence를 기반으로 NMS을 수행해 객체당 하나의 bbox를 남긴다.</li>
  <li>YOLO의 장점은 첫째, YOLO는 detection을 하나의 regression problem으로 취급하기때문에 파이프라인이 복잡하지 않다. 속도가 매우 빠르면서 다른 실시간 탐지 시스템에 비해 두배 이상의 mAP를 기록했다.</li>
  <li>둘째, YOLO는 sliding window나 region proposals 기법과 달리 훈련 및 테스트 시간 동안 고정된 grid cell로 전체 이미지를 보기때문에 클래스에 대한 context 정보나 class의 형태 등과 같은 정보를 암시적으로 인코딩할 수 있다. Fast R-CNN의 경우 large context 정보가 제대로 반영되지 않기때문에 background patches에 예측력이 좋지않다. YOLO는 Fast R-CNN의 절반 정도의 background 예측 에러를 갖고있다.</li>
  <li>셋째, YOLO는 natural한 이미지를 학습했고, 예술 이미지에 대해 test했을 때 일반화 능력이 뛰어났다. DPM, Fast R-CNN보다 새로운 도메인에 대해서도 예측을 잘 수행했다.</li>
  <li>속도가 매우 빠르지만 아직 localization과 작은 객체 탐지에 대해 어려움을 겪고 있고 여러 실험을 통해 속도와 정확도간의 trade-off에 대해 살펴본다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/8cb69ee4-8331-4954-9cae-311d0cd17261" />
  </p>
</div>

<p><br /><br /></p>

<h2 id="unified-detection">Unified Detection</h2>
<p>YOLO network는 한 이미지에 존재하는 모든 클래스에 걸쳐 모든 bbox를 동시에 예측한다. 즉, 네트워크는 전체 이미지와 이미지에 포함된 객체를 전역적으로 예측한다는 것이다. YOLO는 이미지를 SxS의 그리드로 나누는데 만약 객체의 중심이 그리드셀 내에 있다면 해당 그리드셀에서는 반드시 객체가 탐지되어야 할 것이다. 각각의 그리드 셀마다 B개의 bbox와 bbox 내의 confidence score(물체 존재 확신도)를 예측한다. confidence를 수식화하면 $Pr(Object) * IOU(truth|pred)$로 표현한다. 만약, 그리드셀 내에 객체가 없다면 confidence score는 0이 되어야한다. 물체가 존재한다면 confidence score는 IOU와 같아지길 원한다.</p>

<p>Bounding box는 x,y,w,h,confidence score 5개의 예측값을 가진다. x,y는 grid cell 내에서의 중심점 좌표고 w,h는 전체 이미지 내의 넓이와 높이다. Confidence score는 예측 box와 어떤 ground-truth box간의 IOU를 대변한다.</p>

<p>또한, 각각의 그리드 셀은 C(conditonal class probabilities)라는 클래스 확률을 예측한다.</p>

\[C = Pr(Class_{i}|Object)\]

<p>이 조건부 확률은 그리드 셀안에 객체가 있다는 조건 하에 해당 객체가 어떤 class인지에 대한 확률이다. 그리드 셀 내의 Bounding box의 갯수와 상관없이 한 그리드 셀마다 하나의 class probabilities만 예측한다.</p>

<p>Test 단계에서는 C와 cofidence score를 곱해 box내에 특정 class가 존재할 확률과 bounding box가 얼마나 이 객체에 맞게 잘 형성되었는지를 파악한다.</p>

\[Pr(Class_i|Object) * Pr(Object) * IOU(truth|pred) = Pr(Class_i) * IOU(truth|pred)\]

<p>아래 Figure 2.에서 YOLO가 어떻게 detection을 수행하는지 잘 설명했다. Box당 4개의 좌표와 물체 존재 확신도, class probobilities가 필요하고 이것을 그리드마다 수행하게되므로 예측에는 총 S x S x (5B + C)의 텐서가 필요하다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/7141e759-c21c-408e-abe5-d8372502f2b4" />
  </p>
</div>

<p><br /><br /></p>

<h3 id="1-network-design">1. Network Design</h3>
<ul>
  <li>PASCAL VOC 데이터셋에서 평가를 진행한다.</li>
  <li>CNN에서 이미지의 특징을 추출하고 FC층에서 class 확률과 bbox 좌표를 예측한다.</li>
  <li>ImageNet 데이터를 활용하여 classification task로 pretrain했고, detection에서는 224x224크기의 이미지를 두배로 늘려서 사용했다.</li>
  <li>아키텍처는 GoogLeNet과 유사한데 인셉션 모듈대신 1x1, 3x3 합성곱을 사용했다(DarkNet).</li>
  <li>Fast YOLO는 24 layer 대신 9 layer만 사용해서 속도를 빠르게 했다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/616adf0e-4dd2-481f-ab47-a25f97697bd8" />
  </p>
</div>

<p><br /></p>

<h3 id="2-training">2. Training</h3>
<ul>
  <li>사전훈련한 모델은 ImageNet 2012에서 GoogLeNet과 비슷한 수준의 성능을 기록했고, 모든 훈련과 추론에 DarkNet 프레임워크를 사용했다.</li>
  <li>객체 탐지를 수행하기 위해 4개의 convolutional layer와 랜덤하게 가중치가 초기화된 2개의 완전결합층을 추가해서 모델을 조금 수정했다. 또한, detection에는 더 세밀한 작업이 필요해 resolution을 2배로 높였다(448 x 448).</li>
  <li>최종 layer에서는 class 확률과 bbox 좌표를 예측하는데 bbox는 w,h를 정규화해서 0과 1사이가 되도록했고, x,y 좌표도 특정 그리드 셀 위의 오프셋으로 변환하여 0과 1사이로 변환했다.</li>
  <li>최종 layer만 linear activation을 사용했고 모든 다른 layer에서는 Leaky ReLU를 사용했다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="300" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/3210764f-0ce3-48a4-acce-6816e460c4fe" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>MSE를 사용하면, 손실함수가 간단하지만 mAP를 극대화하려는 목표와는 거리가 있다. 또한, grid cell에 물체가 포함되지 않은 경우가 많다면 물체가 있는데도 confidence score가 0으로 되어 예측 자체가
수행되지 않은 경우가 생길 수 있기때문에 학습이 불안정할 수도 있다.</li>
  <li>이를 해결하기 위해, bbox 좌표 예측의 손실을 높이고, 객체가 포함되지 않은 box에 대한 confidence의 예측 손실을 줄이는 방법을 사용했다. localization과 classification 중 localization의 가중치를 더 증가시키고 객체가 없는 confidence loss의 가중치를 있는 가중치보다 더 감소시키는 방법이다(객체가 없는 경우가 훨씬 많기때문에 객체가 있을 때의 loss가 더 중요함). 이는 $\lambda_{coord} = 5$, $\lambda_{noobj} = 0.5$를 설정하여 해결할 수 있다(coordinate는 가중치가 정수라서 원래보다 커지고, noobj는 소수라서 작아짐).</li>
  <li>MSE는 bbox가 크든, 작든 동일한 가중치를 부여한다는 단점이 있다. 오차의 관점에서 크기가 큰 bbox는 작은 bbox보다 오차에 많은 영향을 미친다. 크기가 큰만큼 오차가 상대적으로 더 클 것이기때문이다. 또한, 작은 bbox는 큰 bbox보다 편차에 민감하다. 예를 들어, 큰 객체를 감싸고 있는 bbox가 0.5 움직여도 여전히 객체를 감싸고 있을 수도 있지만, 작은 객체를 감싸고 있는 bbox가 0.5 움직이면 객체에서 벗어날 수도 있기때문이다. 따라서 큰 bbox의 loss 때문에 작은 bbox가 영향을 받는다면 최적화가 잘 이루어지지 않을 수 있다. 이 문제를 해결하기 위해 넓이와 높이를 직접적으로 예측하지 않고, 제곱근을 예측하는 방법을 사용했다. 제곱근을 예측하면 box의 크기가 클수록 증가율이 낮아지기때문에 bbox가 크더라도 제곱근으로 인해 작은 bbox가 큰 영향을 받지 않게되고 loss를 줄일 수 있다. 결과적으로 Box가 클수록 증가율이 작아져 IOU에 적은 영향을 끼치게된다.</li>
  <li>위의 손실함수 관련 내용을 수식으로 표현하면 다음과 같다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/488823e2-1e29-445f-a8f1-07ce68f27c5b" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>$1_{ij}^{obj}$는 grid cell내에 class가 존재한다는 것을 의미하고 존재하면 1, 아니면 0으로 표현한다.</li>
  <li>$1_{ij}^{noobj}$는 객체가 존재하지 않을 때 confidence score loss를 계산하기 위해 사용된다.</li>
  <li>$1_{ij}^{obj}$ 는 grid cell $i$의 $j$번째 bbox predictor가 사용되는지의 여부이다. 위의 수식을 5단계로 표현하면
    <ul>
      <li>먼저, 객체가 존재하고 그리드 셀 $i$의 bbox predictor $j$에 대해 x,y loss를 구한다.</li>
      <li>두번째, 객체가 존재하고 그리드 셀 $i$의 bbox predictor $j$에 대해 w,h loss를 구한다.(큰 bbox의 증가율이 커지지 않도록 제곱근을 예측)</li>
      <li>세번째, 객체가 존재하고 그리드 셀 $i$의 bbox predictor $j$에 대해 confidence loss를 구한다(물체가 존재하기때문에 $C_i=1$)</li>
      <li>네번째, 객체가 존재하지않을때 그리드 셀 $i$의 bbox predictor $j$에 대해 confidence loss를 구한다.($C_i=0$)</li>
      <li>다섯번째, 객체가 존재하고 그리드 셀 $i$의 bbox predictor $j$에 대해 class probabilities loss를 구한다.(class가 맞으면 $p_i(C)=1$ 아니면 0)</li>
      <li>$\lambda_{coord}$ x,y,w,h의 좌표 loss와 다른 loss간의 밸런스를 위한 parameter.</li>
      <li>$\lambda_{noobj}$ 객체가 있는 box와 없는 box의 loss 간의 밸런스를 위한 parameter.</li>
    </ul>
  </li>
  <li>위의 모든 과정을 다 더해서 손실함수를 만들고 모델을 최적화시킨다.</li>
  <li>과적합을 막기위해서 dropout(0.5), 원본 이미지의 크기보다 최대 20%까지 랜덤하게 scaling &amp; translation 적용, HSV factor 조절의 data augmentation을 사용했다.</li>
</ul>

<p><br /></p>

<h3 id="3-inference">3. Inference</h3>
<ul>
  <li>Test 시에는 PASCAL VOC에서 이미지당 98개의 bounding box가 그려졌고 각 박스에 대해 class를 예측했다.</li>
  <li>물체가 너무 크거나 물체가 여러 셀의 경계 근처에 있다면 주변의 다른 셀들의 정보를 참고해야 원활한 localized가 수행될 수 있는데 이러한 경우 하나의 객체가 여러 셀에서 발견되는 multiple detection 문제가 생길 수 있다. NMS를 통해 하나의 객체에 하나의 bounding box만 남도록 했고 RCNN이나 DPM처럼 NMS가 매우 중요하진않지만 이를 사용했을 때 mAP가 조금 향상되었다.</li>
</ul>

<p><br /></p>

<h3 id="4-limitations-of-yolo">4. Limitations of YOLO</h3>
<ul>
  <li>YOLO는 각 그리드 셀마다 두개의 bounding box가 그려지고(각각 다른 종횡비를 가진 bbox를 그리고 NMS기법으로 하나만 남김), 각 셀은 오직 하나의 class로만 예측이되어야하기 떄문에 공간적 제약이 생긴다. 이런 공간적인 제약은 근처에 있는 한 셀에 여러 objects가 있는 경우 모든 objects를 잘 탐지하지 못하고, 특히 크기가 작은 물체를 탐지하는데 어려움을 겪게한다.</li>
  <li>Data로부터 bounding box를 그려내기때문에 가로,세로 비율이나 형태가 익숙하지 않다면 예측에 어려움을 겪는다. 또한, 여러 층을 거치며 downsampling된 features를 사용하기때문에 bbox를 예측하는 단계에서는 input image의 정보가 많이 선명하진 않을 것이다.</li>
  <li>Detection performance를 위해 loss function을 정의했지만, 제곱근을 사용했더라도 작은 bbox나 큰 bbox의 loss function을 결국 유사하게 가져갔고 그 결과 큰 bbox보다 작은 bbox가 IOU에 많은 영향을 미쳤다.</li>
  <li>Error의 가장 큰 문제는 localization이다.</li>
</ul>

<p><br /><br /></p>

<h2 id="comparison-to-other-detection-system">Comparison to Other Detection System</h2>
<p>YOLO detection system이 다른 systems들과 어떤 공통점 혹은 차별점을 가지는지 살펴보자.</p>

<p><strong>Deformable parts models</strong></p>
<ul>
  <li>DPM은 sliding window 개념을 활용하여 객체를 탐지한다.</li>
  <li>DPM은 각각 분리된 형태로 파이프라인을 구축해 features를 추출하고, regions에서 분류를 수행하고, 가장 점수가 높은 regions에서 bbox를 예측한다.</li>
  <li>YOLO는 이 모든 과정을 하나로 통합했다는 점에서 DPM과 차이가 있다. 통합된 아키텍처로 DPM보다 더 빠르고 정확한 모델을 만들었다.</li>
</ul>

<p><br /></p>

<p><strong>R-CNN</strong></p>
<ul>
  <li>R-CNN은 sliding window 대신 selective search 알고리즘으로 다양한 region proposals을 생성하고, CNN으로 features를 추출, SVM으로 분류를 수행, linear model로 bbox 예측, 중복된 bbox는 NMS로 제거하는 기법들을 사용했다.</li>
  <li>굉장히 복잡한 파이프라인이고 각각의 tasks가 모두 독립적으로 이루어져 results를 출력하는 속도가 매우 느리다는 단점을 가지고있다.</li>
  <li>YOLO는 potential bbox를 제안받고, CNN을 통해 features를 추출한다는 점이 동일하지만 selective search가 아닌 공간적 제약을 가진 grid cell로 region proposals을 수행한다는 점에서 차이가 존재한다.</li>
  <li>또한, 2000개의 bbox가 생성되는 R-CNN에 비해 98개의 bbox만 생성되고, 각각의 components를 하나로 통합했다는 점에서 R-CNN과 차이가 있다.</li>
</ul>

<p><br /></p>

<p><strong>Other Fast Detectors</strong></p>
<ul>
  <li>DPM과 R-CNN 모두 각각의 components를 개선시켜 속도와 성능을 높였지만, YOLO는 여전히 애초에 하나의 pipeline에서 속도가 빠른 네트워크를 구축했다는 점에서 차별점이 존재한다.</li>
  <li>YOLO는 general purpose detector로 다양한 객체들을 동시에 예측할 수 있다.</li>
  <li>이외에도 YOLO는 다른 여러가지 detector system보다 빠르고 한 이미지 내에서 single 뿐만 아니라 multiple objects도 잘 탐지할 수 있는 모델이다.</li>
</ul>

<p><br /><br /></p>

<h2 id="experiments">Experiments</h2>
<p>YOLO와 다른 real-time detection systems를 비교하기 위해 VOC 2007 데이터셋을 Fast R-CNN과 비교했다. Error를 계산하는 방법이 다르기때문에 Fast R-CNN을 rescore했고, background false positives 에러를 감소시켜 기존보다 Fast R-CNN의 성능을 높게 조정했다. 또한, VOC 2012의 SOTA mAP와 비교했고, 최종적으로 YOLO의 일반화능력을 평가하기 위해 새로운 도메인인 예술작품에서 다른 시스템과 비교해보았다.</p>

<p><br /></p>

<h3 id="1-comparison-to-other-real-time-systems">1. Comparison to Other Real-Time Systems</h3>
<ul>
  <li>많은 연구들이 networks를 빠르게 만드는데 초점을 맞추고있다. 하지만 실제로 초당 30프레임 이상으로 실행되는 시스템은 DPM 모델밖에 없다. 30Hz or 100Hz로 실행되는 DPM의 GPU와 YOLO를 비교했다.</li>
  <li>Fast YOLO는 현존하는 가장 빠른 detection model이고 52.7%의 mAP로 이전보다 2배 이상 정확한 모델이다. YOLO 또한, mAP 63.4%까지 성능을 끌어올렸다.</li>
  <li>VGG-16을 사용해서 YOLO를 훈련시키면 모델이 정확하지만 속도가 기존보다 떨어졌다. 이는 VGG를 기반으로 한 다른 모델들과 비교하기에는 유용하지만 YOLO보다 많이 느려서 VGG로 훈련시키지 않은 원래 YOLO로 비교를 진행했다.</li>
  <li>DPM은 mAP의 큰 희생없이 속도를 효과적으로 높였으나 여전히 성능은 2배 이상 떨어진다. 특히, 딥러닝을 통한 접근에도 성능이 높지 않다는 점에서 한계를 가지고있다.</li>
  <li>R-CNN에서 R을 빼면 selective search가 static bounding box proposals로 바껴서 R-CNN보다 훨씬 빠른 모델이 만들어진다. R-CNN은 여전히 region proposals에 의존하며 좋은 proposals이 없다면 높은 정확도를 낼 수 없다.</li>
  <li>Fast R-CNN은 R-CNN보다 속도가 빨라졌지만 여전히 selective search에 의존하며 이미지당 2초의 proposals 시간이 소요된다. 따라서, mAP는 높지만 0.5fps라서 실시간으로 취급하긴 어렵다.</li>
  <li>가장 최근 모델인 Faster R-CNN은 selective search를 neural network로 대체해서 성능과 속도에 큰 향상을 이루어냈다. 테스트 결과, 가장 정확도가 높은 모델은 7fps, 더 작은 대신 덜 정확한 모델은 18fps로 실행됐다. Faster R-CNN에 다른 여러가지 모델을 훈련시키면 정확도 향상에 비해 속도가 YOLO보다 크게 느려졌다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/07ab98ed-c07f-460c-a5de-bd2d4ff83e2f" />
  </p>
</div>

<p><br /></p>

<h3 id="2-voc-2007-analysis">2. VOC 2007 Analysis</h3>
<ul>
  <li>현존하는 모델 중 PASCAL에서 가장 좋은 성능을 가지고 있는 Fast R-CNN 모델과 성능을 비교했다. Hoiem의 방법론을 사용했고 각 class에 대해 top N predictions(가장 잘맞춘 N개의 class 예측을 평균내어 에러를 비교)을 확인했다. 각각의 prediction은 correct or 다음과 같은 type of error로 구분된다.
    <ul>
      <li>Correct: correct class and IOU &gt; 0.5</li>
      <li>Localization: correct class, 0.1 &lt; IOU &lt; 0.5</li>
      <li>Similar: class is similar, IOU &gt; 0.1</li>
      <li>Other: class is wrong, IOU &gt; 0.1</li>
      <li>Backgruond: IOU &lt; 0.1 for any object</li>
    </ul>
  </li>
  <li>아래 Figure .4를 통해 YOLO는 localization error가 다른 에러를 합친 것보다 더 클만큼 localization에 어려움을 겪고있다는 것을 알 수 있다.</li>
  <li>Fast R-CNN은 localization error는 작지만 backgound에 대해 예측을 잘 수행하지 못하고있다. Background error는 모델이 object라고 예측했는데 실제로는 배경이었던 false positives error이다. 즉, 배경을 제대로 맞추지 못하는 것이다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/e0b98653-2947-4036-bf21-b0497f718a5a" />
  </p>
</div>

<p><br /></p>

<h3 id="3-combining-fast-r-cnn-and-yolo">3. Combining Fast R-CNN and YOLO</h3>
<ul>
  <li>YOLO가 Fast R-CNN보다 배경을 잘 예측하기때문에 이 둘을 조합하여 성능을 향상시켰다. R-CNN이 예측한 bbox에 대해 YOLO가 유사한 박스를 예측하면 YOLO가 미리 지정한 확률과 두 박스 간의 겹침의 정도에 따라 해당 예측에 boost를 준다.</li>
  <li>그 결과, YOLO와 결합했을 때 Fast R-CNN은 기존보다 3.2% 증가한 75%의 mAP를 달성했다.</li>
  <li>이 결합 모델은 각 모델을 개별적으로 실행한 다음 결과를 합치기때문에 YOLO의 속도 이점을 누릴 순 없지만, Fast R-CNN의 원래 속도에 비해 계산 시간이 크게 추가되진 않았다.</li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/c316d60a-8fe9-4db0-95ed-df8dbff25193" />
  </p>
</div>

<p><br /></p>

<h3 id="4-voc-2012-results">4. VOC 2012 Results</h3>
<ul>
  <li>VOC 2012 test sets에서 YOLO는 57.9% mAP를 기록했다. 이는 VGG-16을 사용한 original R-CNN 모델과 비슷한 성능이다.</li>
  <li>병, 모니터 등과 같은 작은 물체를 잘 예측하지 못했지만, 고양이와 기차 등 다른 카테고리에서는 YOLO가 더 높은 성능을 기록했다. Fast R-CNN과 YOLO 결합 모델은 70.7% mAP를 기록하며 최종 5위를 기록했다.</li>
</ul>

<p><br /></p>

<h3 id="5-generalizability-person-detection-in-artwork">5. Generalizability: Person Detection in Artwork</h3>
<ul>
  <li>현실에서는 모델이 접해보지 못한 수많은 데이터가 존재한다. 따라서, YOLO의 일반화 능력을 평가하기 위해 Picasso Dataset과 People-Art Dataset을 사용하여 예술 작품 속에서 사람을 탐지하는 test를 진행했다.</li>
  <li>성능은 사람만 탐지할 것이기때문에 people class에 대한 average precision을 지표로 사용했다. 모든 모델은 VOC 2007의 people 데이터로 학습했고 Picasso model은 VOC 2012로, People-Art 모델은 VOC 2010으로 학습했다.</li>
  <li>R-CNN은 VOC 2007에서 AP가 높았지만, artwork에서는 AP가 크게 떨어졌다. 분류기 단계에서 작은 regions을 보고 좋은 proposals을 수행해야하기때문에 어려울 것이다.</li>
  <li>DPM은 artwork에서도 비슷한 AP를 가졌다. Object에 대해 공간 정보를 잘 간직한 모델이기때문에 R-CNN보다 감소량이 적을 것이다. 하지만, 애초에 AP 자체가 높지않다는 것이 문제다.</li>
  <li>YOLO는 VOC 2007에서 성능이 우수할 뿐만 아니라 artwork에서도 성능 저하가 작다. YOLO도 DPM처럼 객체의 모양이나 공간 정보를 잘 간직하고, artwork는 일반 이미지와 다르지만 객체의 크기와 모양이 비슷했기때문에 좋은 성능을 냈을 것으로 추측한다.</li>
</ul>

<p><br /><br /></p>

<h2 id="real-time-detection-in-the-wild">Real-Time Detection In The Wild</h2>
<ul>
  <li>웹캠을 통해 야생에서 객체 탐지를 수행했을 때 YOLO는 객체가 움직이거나 모양이 변할 때 이를 감지했다.</li>
</ul>

<p><br /><br /></p>

<h1 id="conclusion">Conclusion</h1>
<ul>
  <li>YOLO는 전체 이미지를 직접적으로 학습하고 간단하게 구축할 수 있는 모델이다.</li>
  <li>Classifier-based approaches와 달리 detection performance에 알맞은 손실함수를 정의했고 모델에 이를 사용했다.</li>
  <li>Fast YOLO는 가장 빠른 general-purpose object detector이고, YOLO는 real-time 객체 탐지에서 SOTA를 기록했다.</li>
  <li>새로운 도메인에 대해서도 빠르면서 일반화가 잘된 detection이 가능하다.</li>
</ul>

<p><br /><br /></p>

<h1 id="개인적인-생각">개인적인 생각</h1>
<ul>
  <li>YOLO는 class classification과 bounding box regression을 단일 네트워크로 구축했다는 점에서 굉장히 의미있는 연구였다.</li>
  <li>다른 논문들과는 달리 기존과는 전혀 다른 도메인인 artwork를 사용해서 일반화 능력을 평가했기때문에 YOLO의 일반화 능력에 더 신뢰가 갔다.</li>
  <li>여러가지 detection model들과 성능을 비교하고 심지어 결합하기까지 해본 실험을 통해 YOLO가 빠르고 성능도 준수한 모델 개발을 위해 얼마나 많은 시간을 들였는지 느껴졌다.</li>
  <li>2023년 기준, YOLOv8이 나와 객체 탐지에서 매우 좋은 성능을 내고있는데 향후 논문들에서는 기존 YOLO가 가지고 있던 localization과 작은 객체 탐지에 대한 문제를 어떻게 해결할지 기대가 된다.</li>
  <li>YOLO에서는 한 그리드셀에서 하나의 classification만 수행이 되도록했는데 성능을 높이기위해서는 그리드셀 내에서도 multiple detection이 가능해야한다. 이와같은 문제를 어떻게 해결할지 궁금하다.</li>
</ul>

<p><br /><br /></p>

<h1 id="구현">구현</h1>
<p>Pytorch로 YOLOv1 모델을 구현해보자(<a href="https://github.com/aladdinpersson/Machine-Learning-Collection/blob/master/ML/Pytorch/object_detection/YOLO/train.py">참고</a>).</p>

<h2 id="1-model">1. Model</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 튜플이면 해당 layer가 하나인 것, list이면 해당 layer가 여러개인 것
</span>
<span class="n">architecture_config</span> <span class="o">=</span> <span class="p">[</span>
    <span class="p">(</span><span class="mi">7</span><span class="p">,</span> <span class="mi">64</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">3</span><span class="p">),</span> <span class="c1"># backbone에서 224 이미지를 2배 키워서 448
</span>    <span class="s">"M"</span><span class="p">,</span> <span class="c1"># max pooing strdie 2, kernel 2 -&gt; 112
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">192</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="c1"># 112
</span>    <span class="s">"M"</span><span class="p">,</span> <span class="c1"># 56
</span>    <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">128</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="c1"># 56
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">256</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="c1"># 56
</span>    <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="c1"># 56
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">512</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="c1"># 56
</span>    <span class="s">"M"</span><span class="p">,</span> <span class="c1"># 28
</span>    <span class="p">[(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">512</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="mi">4</span><span class="p">],</span> <span class="c1"># 28
</span>    <span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">512</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="c1"># 28
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</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="c1"># 28
</span>    <span class="s">"M"</span><span class="p">,</span> <span class="c1"># 14
</span>    <span class="p">[(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">512</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</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="mi">2</span><span class="p">],</span> <span class="c1"># 14
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</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="c1"># 14
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">1</span><span class="p">),</span> <span class="c1"># 7
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</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="c1"># 7
</span>    <span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="mi">1024</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="c1"># 7
</span><span class="p">]</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">CONVBLOCK</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">CONVBLOCK</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">conv</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">bias</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">batchnorm</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">leakyrelu</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">LeakyReLU</span><span class="p">(</span><span class="mf">0.1</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">leakyrelu</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">batchnorm</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">conv</span><span class="p">(</span><span class="n">x</span><span class="p">)))</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">YOLOv1</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">in_channels</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">YOLOv1</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">architecture</span> <span class="o">=</span> <span class="n">architecture_config</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span> <span class="o">=</span> <span class="n">in_channels</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">darknet</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_create_conv_layers</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">architecture</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">fc</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_create_fc</span><span class="p">(</span><span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span><span class="n">x</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">darknet</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">start_dim</span><span class="o">=</span> <span class="mi">1</span><span class="p">)</span>
        <span class="k">return</span> <span class="bp">self</span><span class="p">.</span><span class="n">fc</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">_create_conv_layers</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">architecture</span><span class="p">):</span>
        <span class="n">layers</span> <span class="o">=</span> <span class="p">[]</span>
        <span class="n">in_channles</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span>

        <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">architecture</span><span class="p">:</span>
            <span class="k">if</span> <span class="nb">type</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">==</span> <span class="nb">tuple</span><span class="p">:</span>
                <span class="n">layers</span> <span class="o">+=</span> <span class="p">[</span>
                    <span class="n">CONVBLOCK</span><span class="p">(</span><span class="n">in_channles</span><span class="p">,</span> <span class="n">out_channels</span><span class="o">=</span><span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">kernel_size</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="n">stride</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">padding</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span>
                    <span class="p">)</span>
                <span class="p">]</span>
                <span class="n">in_channles</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>

            <span class="k">elif</span> <span class="nb">type</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">==</span> <span class="nb">str</span><span class="p">:</span>
                <span class="n">layers</span> <span class="o">+=</span> <span class="p">[</span><span class="n">nn</span><span class="p">.</span><span class="n">MaxPool2d</span><span class="p">(</span><span class="n">kernel_size</span><span class="o">=</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span><span class="mi">2</span><span class="p">),</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)]</span>

            <span class="k">elif</span> <span class="nb">type</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">==</span> <span class="nb">list</span><span class="p">:</span>
                <span class="n">conv1</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="n">conv2</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
                <span class="n">num_repeats</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span>

                <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_repeats</span><span class="p">):</span>
                    <span class="n">layers</span> <span class="o">+=</span> <span class="p">[</span>
                        <span class="n">CONVBLOCK</span><span class="p">(</span><span class="n">in_channles</span><span class="p">,</span> <span class="n">out_channels</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">kernel_size</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">stride</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span>
                                  <span class="n">padding</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span>
                                  <span class="p">)</span>
                    <span class="p">]</span>

                    <span class="n">layers</span> <span class="o">+=</span> <span class="p">[</span> <span class="c1"># conv1의 output이 conv2의 in
</span>                        <span class="n">CONVBLOCK</span><span class="p">(</span><span class="n">in_channels</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">out_channels</span> <span class="o">=</span> <span class="n">conv2</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">kernel_size</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">stride</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span>
                                  <span class="n">padding</span> <span class="o">=</span> <span class="n">conv1</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span>
                                  <span class="p">)</span>
                    <span class="p">]</span>
                    <span class="n">in_channles</span> <span class="o">=</span> <span class="n">conv2</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
        <span class="k">return</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span><span class="o">*</span><span class="n">layers</span><span class="p">)</span>

    <span class="k">def</span> <span class="nf">_create_fc</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">num_split_cell</span><span class="p">,</span> <span class="n">num_boxes</span><span class="p">,</span> <span class="n">num_classes</span><span class="p">):</span>
        <span class="n">S</span><span class="p">,</span> <span class="n">B</span><span class="p">,</span> <span class="n">C</span> <span class="o">=</span> <span class="n">num_split_cell</span><span class="p">,</span> <span class="n">num_boxes</span><span class="p">,</span> <span class="n">num_classes</span>

        <span class="n">fc_layer</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Flatten</span><span class="p">(),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">1024</span> <span class="o">*</span> <span class="n">S</span> <span class="o">*</span> <span class="n">S</span><span class="p">,</span> <span class="mi">4096</span><span class="p">),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Dropout</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">LeakyReLU</span><span class="p">(</span><span class="mf">0.1</span><span class="p">),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">4096</span><span class="p">,</span> <span class="n">S</span> <span class="o">*</span> <span class="n">S</span> <span class="o">*</span> <span class="p">(</span><span class="mi">5</span> <span class="o">*</span> <span class="n">B</span> <span class="o">+</span> <span class="n">C</span><span class="p">))</span> <span class="c1"># 각 박스마다 4개좌표 + confidence, 각 class에 속할 확률 값
</span>        <span class="p">)</span>
        <span class="k">return</span> <span class="n">fc_layer</span>
</code></pre></div></div>

<h2 id="2-dataset">2. Dataset</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">from</span> <span class="nn">pycocotools.coco</span> <span class="kn">import</span> <span class="n">COCO</span>
<span class="kn">import</span> <span class="nn">os</span>
<span class="kn">import</span> <span class="nn">cv2</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">transforms</span> <span class="k">as</span> <span class="n">T</span>
<span class="kn">from</span> <span class="nn">PIL</span> <span class="kn">import</span> <span class="n">Image</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>

<span class="k">class</span> <span class="nc">SoccerDataset</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">utils</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="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_path</span><span class="p">,</span> <span class="n">label_path</span><span class="p">,</span> <span class="n">S</span> <span class="o">=</span> <span class="mi">7</span><span class="p">,</span> <span class="n">B</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="n">C</span> <span class="o">=</span> <span class="mi">3</span><span class="p">,</span> <span class="n">transforms</span> <span class="o">=</span> <span class="bp">None</span><span class="p">):</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">data_path</span> <span class="o">=</span> <span class="n">data_path</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">label_path</span> <span class="o">=</span> <span class="n">label_path</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span> <span class="o">=</span> <span class="n">transforms</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">=</span> <span class="n">S</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">B</span> <span class="o">=</span> <span class="n">B</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">C</span> <span class="o">=</span> <span class="n">C</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">imgs</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">data_path</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.jpg'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">labs</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">label_path</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.txt'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>

    <span class="k">def</span> <span class="nf">__getitem__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">idx</span><span class="p">):</span>
        <span class="n">lab</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span>
        <span class="n">boxes</span> <span class="o">=</span> <span class="p">[]</span>
        <span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">label_path</span> <span class="o">+</span> <span class="n">lab</span><span class="p">)</span> <span class="k">as</span> <span class="n">f</span><span class="p">:</span>
            <span class="k">for</span> <span class="n">label</span> <span class="ow">in</span> <span class="n">f</span><span class="p">.</span><span class="n">readlines</span><span class="p">():</span>
                <span class="n">class_label</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">width</span><span class="p">,</span> <span class="n">height</span> <span class="o">=</span> <span class="p">[</span>
                    <span class="nb">float</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="k">if</span> <span class="nb">float</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">!=</span> <span class="nb">int</span><span class="p">(</span><span class="nb">float</span><span class="p">(</span><span class="n">x</span><span class="p">))</span> <span class="k">else</span> <span class="nb">int</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
                    <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">label</span><span class="p">.</span><span class="n">replace</span><span class="p">(</span><span class="s">"</span><span class="se">\n</span><span class="s">"</span><span class="p">,</span> <span class="s">""</span><span class="p">).</span><span class="n">split</span><span class="p">()</span>
                <span class="p">]</span>

                <span class="n">boxes</span><span class="p">.</span><span class="n">append</span><span class="p">([</span><span class="n">class_label</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">width</span><span class="p">,</span> <span class="n">height</span><span class="p">])</span>

        <span class="n">boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">boxes</span><span class="p">)</span>

        <span class="n">img_path</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">data_path</span> <span class="o">+</span> <span class="n">lab</span><span class="p">.</span><span class="n">split</span><span class="p">(</span><span class="s">'.txt'</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span> <span class="o">+</span> <span class="s">'.jpg'</span>
        <span class="c1"># image = cv2.imread(img_path)
</span>        <span class="c1"># image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
</span>        <span class="n">image</span> <span class="o">=</span> <span class="n">Image</span><span class="p">.</span><span class="nb">open</span><span class="p">(</span><span class="n">img_path</span><span class="p">)</span>

        <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">image</span><span class="p">,</span> <span class="n">boxes</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span><span class="p">(</span><span class="n">image</span><span class="p">,</span> <span class="n">boxes</span><span class="p">)</span>

        <span class="c1"># 해당 box가 grid cell에서 몇번째 cell에 속하는지
</span>        <span class="n">label_matrix</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="bp">self</span><span class="p">.</span><span class="n">S</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">C</span> <span class="o">+</span> <span class="p">(</span><span class="mi">5</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">B</span><span class="p">)))</span>
        <span class="k">for</span> <span class="n">box</span> <span class="ow">in</span> <span class="n">boxes</span><span class="p">:</span>
            <span class="n">class_label</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">w</span><span class="p">,</span> <span class="n">h</span> <span class="o">=</span> <span class="n">box</span><span class="p">.</span><span class="n">tolist</span><span class="p">()</span>
            <span class="n">class_label</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">class_label</span><span class="p">)</span>

            <span class="c1"># i,j는 셀의 행과 열을 의미함
</span>            <span class="c1"># y가 세로로 이동하니까 행
</span>            <span class="n">i</span><span class="p">,</span><span class="n">j</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">*</span> <span class="n">y</span><span class="p">),</span> <span class="nb">int</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">*</span> <span class="n">x</span><span class="p">)</span> <span class="c1"># box가 속해있는 셀의 위치
</span>            <span class="n">x_cell</span><span class="p">,</span> <span class="n">y_cell</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">*</span> <span class="n">x</span> <span class="o">-</span> <span class="n">j</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">*</span> <span class="n">y</span> <span class="o">-</span> <span class="n">i</span> <span class="c1"># box가 해당 cell에서 어느 위치에 있는지 파악
</span>            <span class="n">w_cell</span><span class="p">,</span> <span class="n">h_cell</span> <span class="o">=</span> <span class="p">(</span><span class="n">w</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span><span class="p">,</span> <span class="n">h</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span><span class="p">)</span>

            <span class="c1"># 한 그리드 당 하나의 object는 무조건 존재해야하기때문에
</span>            <span class="k">if</span> <span class="n">label_matrix</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">j</span><span class="p">,</span><span class="mi">3</span><span class="p">]</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># 내 클래스가 3개니까 0,1,2 각각 클래스에 대해 존재하는지 여부 다음 3번째 인덱스에서 confidence score.
</span>                <span class="n">label_matrix</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">j</span><span class="p">,</span><span class="mi">3</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span> <span class="c1"># object가 존재한다고 1로 표현
</span>
                <span class="n">box_coord</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="n">x_cell</span><span class="p">,</span> <span class="n">y_cell</span><span class="p">,</span> <span class="n">w_cell</span><span class="p">,</span> <span class="n">h_cell</span><span class="p">])</span>
                <span class="n">label_matrix</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">j</span><span class="p">,</span><span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">]</span> <span class="o">=</span> <span class="n">box_coord</span> <span class="c1"># 좌표 정보도 추가
</span>                <span class="n">label_matrix</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">j</span><span class="p">,</span><span class="n">class_label</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span> <span class="c1"># 해당 class 라벨이 존재하면 1
</span>
        <span class="k">return</span> <span class="n">image</span><span class="p">,</span> <span class="n">label_matrix</span>

    <span class="k">def</span> <span class="nf">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">labs</span><span class="p">)</span> <span class="c1"># label이 있는 이미지만
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torchvision.transforms</span> <span class="k">as</span> <span class="n">transforms</span> <span class="c1"># torchvision 내에 있는 transforms은 PIL 이미지를 받아서 전처리 수행
</span>
<span class="k">class</span> <span class="nc">Compose</span><span class="p">(</span><span class="nb">object</span><span class="p">):</span>
    <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">transforms</span><span class="p">):</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span> <span class="o">=</span> <span class="n">transforms</span>

    <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">img</span><span class="p">,</span> <span class="n">bboxes</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">transforms</span><span class="p">:</span> <span class="c1"># bboxes는 448로 resize and 정규화
</span>            <span class="n">img</span><span class="p">,</span> <span class="n">bboxes</span> <span class="o">=</span> <span class="n">t</span><span class="p">(</span><span class="n">img</span><span class="p">),</span> <span class="n">bboxes</span>

        <span class="k">return</span> <span class="n">img</span><span class="p">,</span> <span class="n">bboxes</span>

<span class="k">def</span> <span class="nf">get_transform</span><span class="p">():</span>
    <span class="n">transform</span> <span class="o">=</span> <span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="p">.</span><span class="n">Resize</span><span class="p">((</span><span class="mi">448</span><span class="p">,</span> <span class="mi">448</span><span class="p">)),</span> <span class="n">transforms</span><span class="p">.</span><span class="n">ToTensor</span><span class="p">(),])</span>
    <span class="k">return</span> <span class="n">transform</span>
</code></pre></div></div>

<h2 id="3-utils">3. Utils</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>

<span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">import</span> <span class="nn">matplotlib.patches</span> <span class="k">as</span> <span class="n">patches</span>
<span class="kn">from</span> <span class="nn">collections</span> <span class="kn">import</span> <span class="n">Counter</span>

<span class="k">def</span> <span class="nf">intersection_over_union</span><span class="p">(</span><span class="n">boxes_preds</span><span class="p">,</span> <span class="n">boxes_labels</span><span class="p">,</span> <span class="n">box_format</span><span class="o">=</span><span class="s">"midpoint"</span><span class="p">):</span>

    <span class="k">if</span> <span class="n">box_format</span> <span class="o">==</span> <span class="s">"midpoint"</span><span class="p">:</span> <span class="c1">#  x_center, y_center, w, h의 포멧일 때
</span>        <span class="n">box1_x1</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box1_y1</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box1_x2</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span> <span class="o">+</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box1_y2</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span> <span class="o">+</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box2_x1</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box2_y1</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span> <span class="o">-</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box2_x2</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span> <span class="o">+</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>
        <span class="n">box2_y2</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span> <span class="o">+</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">/</span> <span class="mi">2</span>

    <span class="k">if</span> <span class="n">box_format</span> <span class="o">==</span> <span class="s">"corners"</span><span class="p">:</span> <span class="c1"># x1,y1,x2,y2의 포멧일 때
</span>        <span class="n">box1_x1</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span>
        <span class="n">box1_y1</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span>
        <span class="n">box1_x2</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span>
        <span class="n">box1_y2</span> <span class="o">=</span> <span class="n">boxes_preds</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span>  <span class="c1"># (N, 1)
</span>        <span class="n">box2_x1</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">:</span><span class="mi">1</span><span class="p">]</span>
        <span class="n">box2_y1</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span>
        <span class="n">box2_x2</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span>
        <span class="n">box2_y2</span> <span class="o">=</span> <span class="n">boxes_labels</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span>

    <span class="n">x1</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">box1_x1</span><span class="p">,</span> <span class="n">box2_x1</span><span class="p">)</span>
    <span class="n">y1</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">box1_y1</span><span class="p">,</span> <span class="n">box2_y1</span><span class="p">)</span>
    <span class="n">x2</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">min</span><span class="p">(</span><span class="n">box1_x2</span><span class="p">,</span> <span class="n">box2_x2</span><span class="p">)</span>
    <span class="n">y2</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">min</span><span class="p">(</span><span class="n">box1_y2</span><span class="p">,</span> <span class="n">box2_y2</span><span class="p">)</span>

    <span class="c1"># clamp는 겹치는 부분이 없으면 0으로 취급
</span>    <span class="n">intersection</span> <span class="o">=</span> <span class="p">(</span><span class="n">x2</span> <span class="o">-</span> <span class="n">x1</span><span class="p">).</span><span class="n">clamp</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">y2</span> <span class="o">-</span> <span class="n">y1</span><span class="p">).</span><span class="n">clamp</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>

    <span class="n">box1_area</span> <span class="o">=</span> <span class="nb">abs</span><span class="p">((</span><span class="n">box1_x2</span> <span class="o">-</span> <span class="n">box1_x1</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">box1_y2</span> <span class="o">-</span> <span class="n">box1_y1</span><span class="p">))</span>
    <span class="n">box2_area</span> <span class="o">=</span> <span class="nb">abs</span><span class="p">((</span><span class="n">box2_x2</span> <span class="o">-</span> <span class="n">box2_x1</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">box2_y2</span> <span class="o">-</span> <span class="n">box2_y1</span><span class="p">))</span>

    <span class="k">return</span> <span class="n">intersection</span> <span class="o">/</span> <span class="p">(</span><span class="n">box1_area</span> <span class="o">+</span> <span class="n">box2_area</span> <span class="o">-</span> <span class="n">intersection</span> <span class="o">+</span> <span class="mf">1e-6</span><span class="p">)</span> <span class="c1"># 분모가 0이 되는 것을 방지 하기 위해 1e-6
</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">non_max_suppression</span><span class="p">(</span><span class="n">bboxes</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="p">,</span> <span class="n">threshold</span><span class="p">,</span> <span class="n">box_format</span><span class="o">=</span><span class="s">"corners"</span><span class="p">):</span>
    <span class="c1"># bboxes: [label, prob_score, x1~y2]
</span>    <span class="c1"># threshold: confidence score
</span>
    <span class="k">assert</span> <span class="nb">type</span><span class="p">(</span><span class="n">bboxes</span><span class="p">)</span> <span class="o">==</span> <span class="nb">list</span>

    <span class="n">bboxes</span> <span class="o">=</span> <span class="p">[</span><span class="n">box</span> <span class="k">for</span> <span class="n">box</span> <span class="ow">in</span> <span class="n">bboxes</span> <span class="k">if</span> <span class="n">box</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">&gt;</span> <span class="n">threshold</span><span class="p">]</span> <span class="c1"># conf score로 걸러내기
</span>    <span class="n">bboxes</span> <span class="o">=</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">bboxes</span><span class="p">,</span> <span class="n">key</span><span class="o">=</span><span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">x</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">reverse</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="n">bboxes_after_nms</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="k">while</span> <span class="n">bboxes</span><span class="p">:</span>
        <span class="n">chosen_box</span> <span class="o">=</span> <span class="n">bboxes</span><span class="p">.</span><span class="n">pop</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
        <span class="n">bboxes</span> <span class="o">=</span> <span class="p">[</span>
            <span class="n">box</span> <span class="k">for</span> <span class="n">box</span> <span class="ow">in</span> <span class="n">bboxes</span>
            <span class="k">if</span> <span class="n">box</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">!=</span> <span class="n">chosen_box</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="c1"># 현재 박스와 다른 박스들의 클래스가 다르면 다른 객체를 담고있는 것이기때문에 선택되어야함
</span>            <span class="ow">or</span> <span class="n">intersection_over_union</span><span class="p">(</span>
                <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">chosen_box</span><span class="p">[</span><span class="mi">2</span><span class="p">:]),</span>
                <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">box</span><span class="p">[</span><span class="mi">2</span><span class="p">:]),</span>
                <span class="n">box_format</span><span class="o">=</span><span class="n">box_format</span><span class="p">,</span>
            <span class="p">)</span> <span class="o">&lt;</span> <span class="n">iou_threshold</span> <span class="c1"># 만약 같은 클래스라도 많이 겹치지 않으면 다른 객체를 지정하고 있는 것이기때문에 선택되어야함
</span>        <span class="p">]</span>

        <span class="n">bboxes_after_nms</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">chosen_box</span><span class="p">)</span>

    <span class="k">return</span> <span class="n">bboxes_after_nms</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">mean_average_precision</span><span class="p">(</span><span class="n">pred_boxes</span><span class="p">,</span> <span class="n">true_boxes</span><span class="p">,</span> <span class="n">iou_threshold</span> <span class="o">=</span> <span class="mf">0.5</span><span class="p">,</span>
                           <span class="n">box_format</span> <span class="o">=</span> <span class="s">'midpoint'</span><span class="p">,</span> <span class="n">num_classes</span> <span class="o">=</span> <span class="mi">3</span><span class="p">):</span>
    <span class="c1"># bboxes: [train_idx, label, prob_score, x1, y1, x2, y2]
</span>    <span class="n">average_precisions</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">epsilon</span>  <span class="o">=</span> <span class="mf">1e-6</span>

    <span class="k">for</span> <span class="n">c</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_classes</span><span class="p">):</span>
        <span class="n">detections</span> <span class="o">=</span> <span class="p">[]</span>
        <span class="n">ground_truths</span> <span class="o">=</span> <span class="p">[]</span>

        <span class="k">for</span> <span class="n">detection</span> <span class="ow">in</span> <span class="n">pred_boxes</span><span class="p">:</span>
            <span class="k">if</span> <span class="n">detection</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="n">c</span><span class="p">:</span> <span class="c1"># 클래스를 맞췄으면
</span>                <span class="n">detections</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">detection</span><span class="p">)</span>

        <span class="k">for</span> <span class="n">true_box</span> <span class="ow">in</span> <span class="n">true_boxes</span><span class="p">:</span>
            <span class="k">if</span> <span class="n">true_box</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="n">c</span><span class="p">:</span>
                <span class="n">ground_truths</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">true_box</span><span class="p">)</span>

        <span class="c1"># 0번 이미지에 객체가 3개, 1번에 5개면 amount_bboxes = {0:3, 1:5}
</span>        <span class="n">amount_bboxes</span> <span class="o">=</span> <span class="n">Counter</span><span class="p">([</span><span class="n">gt</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="k">for</span> <span class="n">gt</span> <span class="ow">in</span> <span class="n">ground_truths</span><span class="p">])</span> <span class="c1"># train_idx를 카운트 = 실제로 해당 이미지 내 객체가 몇개있는지 카운트
</span>
        <span class="k">for</span> <span class="n">key</span><span class="p">,</span> <span class="n">val</span> <span class="ow">in</span> <span class="n">amount_bboxes</span><span class="p">.</span><span class="n">items</span><span class="p">():</span>
            <span class="c1"># ammount_bboxes = {0:torch.tensor[0,0,0], 1:torch.tensor[0,0,0,0,0]}
</span>            <span class="n">amount_bboxes</span><span class="p">[</span><span class="n">key</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">val</span><span class="p">)</span>

        <span class="n">detections</span><span class="p">.</span><span class="n">sort</span><span class="p">(</span><span class="n">key</span><span class="o">=</span><span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">reverse</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span> <span class="c1"># conf_score에 따라 정렬
</span>        <span class="n">TP</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="nb">len</span><span class="p">(</span><span class="n">detections</span><span class="p">)))</span>
        <span class="n">FP</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="nb">len</span><span class="p">(</span><span class="n">detections</span><span class="p">)))</span>
        <span class="n">total_true_bboxes</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">ground_truths</span><span class="p">)</span>

        <span class="k">if</span> <span class="n">total_true_bboxes</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># 원래 겍체가 없는 이미지라면 다음 클래스로 넘어가기
</span>            <span class="k">continue</span>

        <span class="k">for</span> <span class="n">detection_idx</span><span class="p">,</span> <span class="n">detection</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">detections</span><span class="p">):</span>
            <span class="n">ground_truth_img</span> <span class="o">=</span> <span class="p">[</span>
                <span class="n">bbox</span> <span class="k">for</span> <span class="n">bbox</span> <span class="ow">in</span> <span class="n">ground_truths</span> <span class="k">if</span> <span class="n">bbox</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">==</span> <span class="n">detection</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
            <span class="p">]</span> <span class="c1"># 해당 이미지에서 예측한 객체가 실제로 몇개인지
</span>            <span class="c1"># 하나의 예측 객체에 대해서 정답이랑 가장 가까운 박스가 best iou가 된다.
</span>
            <span class="n">best_iou</span> <span class="o">=</span> <span class="mi">0</span>

            <span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">gt</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">ground_truth_img</span><span class="p">):</span>
                <span class="n">iou</span> <span class="o">=</span> <span class="n">intersection_over_union</span><span class="p">(</span>
                    <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">detection</span><span class="p">[</span><span class="mi">3</span><span class="p">:]),</span>
                    <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">gt</span><span class="p">[</span><span class="mi">3</span><span class="p">:]),</span>
                    <span class="n">box_format</span><span class="o">=</span><span class="n">box_format</span>
                <span class="p">)</span>

                <span class="k">if</span> <span class="n">iou</span> <span class="o">&gt;</span> <span class="n">best_iou</span><span class="p">:</span>
                    <span class="n">best_iou</span> <span class="o">=</span> <span class="n">iou</span>
                    <span class="n">best_gt_idx</span> <span class="o">=</span> <span class="n">idx</span>

            <span class="k">if</span> <span class="n">best_iou</span> <span class="o">&gt;</span> <span class="n">iou_threshold</span><span class="p">:</span> <span class="c1"># iou가 높으면 맞췄다고 가정
</span>                <span class="k">if</span> <span class="n">amount_bboxes</span><span class="p">[</span><span class="n">detection</span><span class="p">[</span><span class="mi">0</span><span class="p">]][</span><span class="n">best_gt_idx</span><span class="p">]</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
                    <span class="n">TP</span><span class="p">[</span><span class="n">detection_idx</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>
                    <span class="c1"># amount_bboxes -&gt; {0:torch.tensor[0,0,0], 1:torch.tensor[0,0,0,0,0]}
</span>                    <span class="n">amount_bboxes</span><span class="p">[</span><span class="n">detection</span><span class="p">[</span><span class="mi">0</span><span class="p">]][</span><span class="n">best_gt_idx</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>
                <span class="k">else</span><span class="p">:</span> <span class="c1"># 이미 해당 객체를 best로 예측했었는데 또 예측한 거면 FP
</span>                    <span class="n">FP</span><span class="p">[</span><span class="n">detection_idx</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>

            <span class="k">else</span><span class="p">:</span>
                <span class="n">FP</span><span class="p">[</span><span class="n">detection_idx</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>

        <span class="n">TP_cumsum</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cumsum</span><span class="p">(</span><span class="n">TP</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
        <span class="n">FP_cumsum</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cumsum</span><span class="p">(</span><span class="n">FP</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
        <span class="n">recalls</span> <span class="o">=</span> <span class="n">TP_cumsum</span> <span class="o">/</span> <span class="p">(</span><span class="n">total_true_bboxes</span> <span class="o">+</span> <span class="n">epsilon</span><span class="p">)</span>
        <span class="n">precisions</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">divide</span><span class="p">(</span><span class="n">TP_cumsum</span><span class="p">,</span> <span class="p">(</span><span class="n">TP_cumsum</span> <span class="o">+</span> <span class="n">FP_cumsum</span> <span class="o">+</span> <span class="n">epsilon</span><span class="p">))</span>
        <span class="n">precisions</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="mi">1</span><span class="p">]),</span> <span class="n">precisions</span><span class="p">))</span> <span class="c1"># precitions의 시작점 1. y축
</span>        <span class="n">recalls</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">([</span><span class="mi">0</span><span class="p">]),</span> <span class="n">recalls</span><span class="p">))</span><span class="c1"># recall의 시작점 0. x축
</span>
        <span class="n">average_precisions</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">trapz</span><span class="p">(</span><span class="n">precisions</span><span class="p">,</span> <span class="n">recalls</span><span class="p">))</span> <span class="c1"># 각 클래스마다 밑면적 구하기. 적분 수행
</span>
    <span class="k">return</span> <span class="nb">sum</span><span class="p">(</span><span class="n">average_precisions</span><span class="p">)</span> <span class="o">/</span> <span class="nb">len</span><span class="p">(</span><span class="n">average_precisions</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">get_bboxes</span><span class="p">(</span><span class="n">loader</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="p">,</span> <span class="n">threshold</span><span class="p">,</span> <span class="n">pred_format</span> <span class="o">=</span> <span class="s">'cells'</span><span class="p">,</span>
               <span class="n">box_format</span><span class="o">=</span><span class="s">'midpoint'</span><span class="p">,</span> <span class="n">device</span> <span class="o">=</span> <span class="s">'cuda'</span><span class="p">):</span>
    <span class="n">all_pred_boxes</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">all_true_boxes</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
    <span class="n">train_idx</span> <span class="o">=</span> <span class="mi">0</span>

    <span class="k">for</span> <span class="n">batch_idx</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">,</span><span class="n">labels</span><span class="p">)</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">loader</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
        <span class="n">labels</span> <span class="o">=</span> <span class="n">labels</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>

        <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
            <span class="n">predictions</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">batch_size</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
        <span class="n">true_bboxes</span> <span class="o">=</span> <span class="n">cellboxes_to_boxes</span><span class="p">(</span><span class="n">labels</span><span class="p">)</span>
        <span class="n">bboxes</span> <span class="o">=</span> <span class="n">cellboxes_to_boxes</span><span class="p">(</span><span class="n">predictions</span><span class="p">)</span>

        <span class="k">for</span> <span class="n">idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">batch_size</span><span class="p">):</span>
            <span class="n">nms_boxes</span> <span class="o">=</span> <span class="n">non_max_suppression</span><span class="p">(</span>
                <span class="n">bboxes</span><span class="p">[</span><span class="n">idx</span><span class="p">],</span>
                <span class="n">iou_threshold</span><span class="o">=</span><span class="n">iou_threshold</span><span class="p">,</span>
                <span class="n">threshold</span><span class="o">=</span><span class="n">threshold</span><span class="p">,</span>
                <span class="n">box_format</span><span class="o">=</span><span class="n">box_format</span>
            <span class="p">)</span>

            <span class="k">for</span> <span class="n">nms_box</span> <span class="ow">in</span> <span class="n">nms_boxes</span><span class="p">:</span>
                <span class="n">all_pred_boxes</span><span class="p">.</span><span class="n">append</span><span class="p">([</span><span class="n">train_idx</span><span class="p">]</span> <span class="o">+</span> <span class="n">nms_box</span><span class="p">)</span> <span class="c1"># nms box 앞에 train_idx 추가
</span>
            <span class="k">for</span> <span class="n">box</span> <span class="ow">in</span> <span class="n">true_bboxes</span><span class="p">[</span><span class="n">idx</span><span class="p">]:</span>
                <span class="k">if</span> <span class="n">box</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">&gt;</span> <span class="n">threshold</span><span class="p">:</span>
                    <span class="n">all_true_boxes</span><span class="p">.</span><span class="n">append</span><span class="p">([</span><span class="n">train_idx</span><span class="p">]</span> <span class="o">+</span> <span class="n">box</span><span class="p">)</span>

            <span class="n">train_idx</span> <span class="o">+=</span> <span class="mi">1</span>
    <span class="n">model</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
    <span class="k">return</span> <span class="n">all_pred_boxes</span><span class="p">,</span> <span class="n">all_true_boxes</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 그리드 셀의 박스를 전체 이미지에 대한 비율로 다시 바꿈
</span><span class="k">def</span> <span class="nf">convert_cellboxes</span><span class="p">(</span><span class="n">predictions</span><span class="p">,</span> <span class="n">S</span><span class="o">=</span><span class="mi">7</span><span class="p">):</span>
    <span class="n">predictions</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="s">'cpu'</span><span class="p">)</span>
    <span class="n">batch_size</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
    <span class="n">predictions</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">batch_size</span><span class="p">,</span> <span class="n">S</span><span class="p">,</span> <span class="n">S</span><span class="p">,</span> <span class="mi">13</span><span class="p">)</span>
    <span class="n">bboxes1</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">]</span>
    <span class="n">bboxes2</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">13</span><span class="p">]</span>
    <span class="n">scores</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">(</span>
        <span class="p">(</span><span class="n">predictions</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">].</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">8</span><span class="p">].</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)),</span> <span class="n">dim</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="p">)</span>
    <span class="n">best_box</span> <span class="o">=</span> <span class="n">scores</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
    <span class="n">best_boxes</span> <span class="o">=</span> <span class="n">bboxes1</span> <span class="o">*</span> <span class="p">(</span><span class="mi">1</span><span class="o">-</span><span class="n">best_box</span><span class="p">)</span> <span class="o">+</span> <span class="n">best_box</span> <span class="o">*</span> <span class="n">bboxes2</span>

    <span class="c1"># arnage(7): 0부터 6까지 값을 갖는 1차원 텐서
</span>    <span class="c1"># repeat(): 배치사이즈만큼 S번 반복해서 3차원 텐서로
</span>    <span class="c1"># unsueeeze(-1): 차원을 하나 추가해서 4차원으로
</span>    <span class="c1"># cell_ihdices[0,2,3]은 첫번째 이미지에서 (2,3)그리드의 box 좌표
</span>    <span class="n">cell_indices</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">S</span><span class="p">).</span><span class="n">repeat</span><span class="p">(</span><span class="n">batch_size</span><span class="p">,</span> <span class="n">S</span><span class="p">,</span> <span class="mi">1</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
    <span class="n">x</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">/</span> <span class="n">S</span> <span class="o">*</span> <span class="p">(</span><span class="n">best_boxes</span><span class="p">[...,</span> <span class="p">:</span><span class="mi">1</span><span class="p">]</span> <span class="o">+</span> <span class="n">cell_indices</span><span class="p">)</span>
    <span class="n">y</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">/</span> <span class="n">S</span> <span class="o">*</span> <span class="p">(</span><span class="n">best_boxes</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">:</span><span class="mi">2</span><span class="p">]</span> <span class="o">+</span> <span class="n">cell_indices</span><span class="p">.</span><span class="n">permute</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">))</span> <span class="c1"># y값을 기준으로 거리 계산
</span>    <span class="n">w_y</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">/</span> <span class="n">S</span> <span class="o">*</span> <span class="n">best_boxes</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="c1"># 2차원
</span>    <span class="n">converted_bboxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">((</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">w_y</span><span class="p">),</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
    <span class="n">predicted_class</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="p">:</span><span class="mi">3</span><span class="p">].</span><span class="n">argmax</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">).</span><span class="n">unsqueeze</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
    <span class="n">best_confidence</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">predictions</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">],</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">8</span><span class="p">]).</span><span class="n">unsqueeze</span><span class="p">(</span>
        <span class="o">-</span><span class="mi">1</span>
    <span class="p">)</span>
    <span class="n">converted_preds</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">(</span>
        <span class="p">(</span><span class="n">predicted_class</span><span class="p">,</span> <span class="n">best_confidence</span><span class="p">,</span> <span class="n">converted_bboxes</span><span class="p">),</span> <span class="n">dim</span><span class="o">=-</span><span class="mi">1</span>
    <span class="p">)</span>

    <span class="k">return</span> <span class="n">converted_preds</span> <span class="c1"># (batch, S, S, (label,confidence, x, y, w_y)) w_y가 2차원이므로 (batch, S, S, 6)
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 전체 이미지에 대한 비율로 바뀐 boxes들을 각각 개별 바운딩 박스로 저장
</span><span class="k">def</span> <span class="nf">cellboxes_to_boxes</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">S</span><span class="o">=</span><span class="mi">7</span><span class="p">):</span>
    <span class="n">converted_pred</span> <span class="o">=</span> <span class="n">convert_cellboxes</span><span class="p">(</span><span class="n">out</span><span class="p">).</span><span class="n">reshape</span><span class="p">(</span><span class="n">out</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">S</span> <span class="o">*</span> <span class="n">S</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># 7,7, 박스 좌표
</span>    <span class="n">converted_pred</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">]</span> <span class="o">=</span> <span class="n">converted_pred</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">].</span><span class="nb">long</span><span class="p">()</span>
    <span class="n">all_bboxes</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="k">for</span> <span class="n">ex_idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">out</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]):</span>
        <span class="n">bboxes</span> <span class="o">=</span> <span class="p">[]</span>

        <span class="k">for</span> <span class="n">bbox_idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">S</span> <span class="o">*</span> <span class="n">S</span><span class="p">):</span>
            <span class="n">bboxes</span><span class="p">.</span><span class="n">append</span><span class="p">([</span><span class="n">x</span><span class="p">.</span><span class="n">item</span><span class="p">()</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">converted_pred</span><span class="p">[</span><span class="n">ex_idx</span><span class="p">,</span> <span class="n">bbox_idx</span><span class="p">,</span> <span class="p">:]])</span>
        <span class="n">all_bboxes</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">bboxes</span><span class="p">)</span>

    <span class="k">return</span> <span class="n">all_bboxes</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">save_checkpoint</span><span class="p">(</span><span class="n">state</span><span class="p">,</span> <span class="n">filename</span><span class="o">=</span><span class="s">"my_checkpoint.pth.tar"</span><span class="p">):</span>
    <span class="k">print</span><span class="p">(</span><span class="s">"=&gt; Saving checkpoint"</span><span class="p">)</span>
    <span class="n">torch</span><span class="p">.</span><span class="n">save</span><span class="p">(</span><span class="n">state</span><span class="p">,</span> <span class="n">filename</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">load_checkpoint</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">):</span>
    <span class="k">print</span><span class="p">(</span><span class="s">"=&gt; Loading checkpoint"</span><span class="p">)</span>
    <span class="n">model</span><span class="p">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">[</span><span class="s">"state_dict"</span><span class="p">])</span>
    <span class="n">optimizer</span><span class="p">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">[</span><span class="s">"optimizer"</span><span class="p">])</span>
</code></pre></div></div>

<h2 id="4-loss">4. Loss</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>

<span class="k">class</span> <span class="nc">YoloLoss</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">S</span><span class="o">=</span><span class="mi">7</span><span class="p">,</span> <span class="n">B</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">C</span><span class="o">=</span><span class="mi">3</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">YoloLoss</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">mse</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">MSELoss</span><span class="p">(</span><span class="n">reduction</span><span class="o">=</span><span class="s">"sum"</span><span class="p">)</span> <span class="c1"># 기본은 MSE를 따르고
</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">S</span> <span class="o">=</span> <span class="n">S</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">B</span> <span class="o">=</span> <span class="n">B</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">C</span> <span class="o">=</span> <span class="n">C</span>

        <span class="c1"># 논문에서 언급한대로 object가 없을 때 0.5, localization에 대해 5의 람다값 부여
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">lambda_noobj</span> <span class="o">=</span> <span class="mf">0.5</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">lambda_coord</span> <span class="o">=</span> <span class="mi">5</span>

    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">target</span><span class="p">):</span>
        <span class="c1"># predictions은 (BATCH_SIZE, S*S(C+B*5)) 사이즈를 가지도록 reshape
</span>        <span class="c1"># [..., 0 ~ 2] class 확률, 3 box1의 confidence score, 4~7 box1 좌표, 8, box2의 confidence scores, 9~12 box2 좌표
</span>        <span class="n">predictions</span> <span class="o">=</span> <span class="n">predictions</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</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">S</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">S</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">C</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">B</span> <span class="o">*</span> <span class="mi">5</span><span class="p">)</span>

        <span class="c1"># cell당 2개의 box가 그려지기때문에 2개의 iou
</span>        <span class="n">iou_b1</span> <span class="o">=</span> <span class="n">intersection_over_union</span><span class="p">(</span><span class="n">predictions</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">],</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">])</span>
        <span class="n">iou_b2</span> <span class="o">=</span> <span class="n">intersection_over_union</span><span class="p">(</span><span class="n">predictions</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">13</span><span class="p">],</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">])</span> <span class="c1"># target은 동일함
</span>        <span class="n">ious</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">iou_b1</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">iou_b2</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>

        <span class="c1"># best
</span>        <span class="n">iou_maxes</span><span class="p">,</span> <span class="n">bestbox</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">ious</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span> <span class="c1"># 두 박스 중 iou 최대값 리턴, 몇번째 박스의 iou가 최대인지 idx 리턴,
</span>        <span class="n">exists_box</span> <span class="o">=</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">].</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>  <span class="c1"># 해당 박스에 객체가 존재하는지. 존재하면 1, 아니면 0. 논문에서 말한 1obj_i
</span>
        <span class="c1"># 1. localization loss 구하기
</span>        <span class="c1"># 두 박스 중 bestbox를 고르는 방법
</span>        <span class="n">box_predictions</span> <span class="o">=</span> <span class="n">exists_box</span> <span class="o">*</span> <span class="p">(</span>
            <span class="p">(</span>
                <span class="n">bestbox</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">13</span><span class="p">]</span> <span class="c1"># best box가 1이라면 두번째 박스를 골라야하니까 9:13을
</span>                <span class="o">+</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="n">bestbox</span><span class="p">)</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">]</span> <span class="c1"># best box가 0이라면 첫번째 박스를 골라야하니까 4:8을 선택. 굳이 복잡하게 표현
</span>            <span class="p">)</span>
        <span class="p">)</span>

        <span class="n">box_targets</span> <span class="o">=</span> <span class="n">exists_box</span> <span class="o">*</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">]</span>

        <span class="c1"># 제곱근 사용
</span>        <span class="c1"># w,h가 항상 양수가 되도록 sign,abs함수로 조정 후 제곱근.
</span>        <span class="n">box_predictions</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">sign</span><span class="p">(</span><span class="n">box_predictions</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">])</span> <span class="o">*</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="nb">abs</span><span class="p">(</span><span class="n">box_predictions</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">+</span> <span class="mf">1e-6</span><span class="p">)</span>
        <span class="p">)</span>
        <span class="n">box_targets</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">box_targets</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">:</span><span class="mi">4</span><span class="p">])</span>

        <span class="c1"># 제곱근을 취한 두 박스 간의 Loss 계산
</span>        <span class="n">box_loss</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mse</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">box_predictions</span><span class="p">,</span> <span class="n">end_dim</span><span class="o">=-</span><span class="mi">2</span><span class="p">),</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">box_targets</span><span class="p">,</span> <span class="n">end_dim</span><span class="o">=-</span><span class="mi">2</span><span class="p">),</span>
        <span class="p">)</span>

        <span class="c1"># 2. Object loss. confidence score
</span>        <span class="c1"># 위와 같은 방법으로 best box 선정. 9는 box2의 객체 존재여부, 3은 box 1의 객체 존재여부
</span>        <span class="c1"># best box를 선정했으니까 해당 box에 대해서만 confidence score 손실 계산
</span>        <span class="c1"># 두 박스 중 best box가 객체가 있는 박스라고 손실을 계산했으면 나머지 다른 하나는 객체가 없는 박스(no_obj_loss)로 손실을 계산해야함.
</span>        <span class="c1"># 한 그리드당 하나의 객체만 탐지해야하기때문에.
</span>
        <span class="n">pred_box</span> <span class="o">=</span> <span class="p">(</span>
            <span class="n">bestbox</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">10</span><span class="p">]</span> <span class="o">+</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="n">bestbox</span><span class="p">)</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]</span>
        <span class="p">)</span>

        <span class="n">object_loss</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mse</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">exists_box</span> <span class="o">*</span> <span class="n">pred_box</span><span class="p">),</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">exists_box</span> <span class="o">*</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">]),</span> <span class="c1"># 타겟의 idx3은 객체 존재여부
</span>        <span class="p">)</span>

        <span class="c1"># 3. No Object loss
</span>        <span class="c1"># No object에 대해서는 두 박스 모두 계산.
</span>        <span class="c1"># start_dim 1은 첫번째 차원을 제외하고 그 이후 차원 flatten
</span>        <span class="c1"># Object loss와는 별개로 두 박스 모두에 객체가 없을 수 있기때문에 여기서는 두 박스를 모두 사용해 손실 계산.
</span>        <span class="n">no_object_loss</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mse</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">((</span><span class="mi">1</span> <span class="o">-</span> <span class="n">exists_box</span><span class="p">)</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">],</span> <span class="n">start_dim</span><span class="o">=</span><span class="mi">1</span><span class="p">),</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">((</span><span class="mi">1</span> <span class="o">-</span> <span class="n">exists_box</span><span class="p">)</span> <span class="o">*</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">],</span> <span class="n">start_dim</span><span class="o">=</span><span class="mi">1</span><span class="p">),</span>
        <span class="p">)</span>

        <span class="n">no_object_loss</span> <span class="o">+=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mse</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">((</span><span class="mi">1</span> <span class="o">-</span> <span class="n">exists_box</span><span class="p">)</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="mi">8</span><span class="p">:</span><span class="mi">9</span><span class="p">],</span> <span class="n">start_dim</span><span class="o">=</span><span class="mi">1</span><span class="p">),</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">((</span><span class="mi">1</span> <span class="o">-</span> <span class="n">exists_box</span><span class="p">)</span> <span class="o">*</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">:</span><span class="mi">4</span><span class="p">],</span> <span class="n">start_dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
        <span class="p">)</span>

        <span class="c1"># 4. Class loss
</span>        <span class="c1"># 모든 클래스들의 확률값을 사용해서 손실계산
</span>        <span class="n">class_loss</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">mse</span><span class="p">(</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">exists_box</span> <span class="o">*</span> <span class="n">predictions</span><span class="p">[...,</span> <span class="p">:</span><span class="mi">3</span><span class="p">],</span> <span class="n">end_dim</span><span class="o">=-</span><span class="mi">2</span><span class="p">,),</span>
            <span class="n">torch</span><span class="p">.</span><span class="n">flatten</span><span class="p">(</span><span class="n">exists_box</span> <span class="o">*</span> <span class="n">target</span><span class="p">[...,</span> <span class="p">:</span><span class="mi">3</span><span class="p">],</span> <span class="n">end_dim</span><span class="o">=-</span><span class="mi">2</span><span class="p">,),</span>
        <span class="p">)</span>

        <span class="n">loss</span> <span class="o">=</span> <span class="p">(</span>
            <span class="bp">self</span><span class="p">.</span><span class="n">lambda_coord</span> <span class="o">*</span> <span class="n">box_loss</span>  <span class="c1"># 논문 손실함수의 첫번째 두번째 부분
</span>            <span class="o">+</span> <span class="n">object_loss</span>  <span class="c1"># 세번째 부분
</span>            <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">lambda_noobj</span> <span class="o">*</span> <span class="n">no_object_loss</span>  <span class="c1"># 네번째 부분
</span>            <span class="o">+</span> <span class="n">class_loss</span>  <span class="c1"># 다섯번째 부분
</span>        <span class="p">)</span>

        <span class="k">return</span> <span class="n">loss</span>
</code></pre></div></div>

<h2 id="5-training">5. Training</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torchvision.transforms</span> <span class="k">as</span> <span class="n">transforms</span>
<span class="kn">import</span> <span class="nn">torch.optim</span> <span class="k">as</span> <span class="n">optim</span>
<span class="kn">import</span> <span class="nn">torchvision.transforms.functional</span> <span class="k">as</span> <span class="n">FT</span>
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
<span class="kn">from</span> <span class="nn">torch.utils.data</span> <span class="kn">import</span> <span class="n">DataLoader</span>
<span class="c1"># from model import Yolov1
# from dataset import VOCDataset
# from utils import (
#     non_max_suppression,
#     mean_average_precision,
#     intersection_over_union,
#     cellboxes_to_boxes,
#     get_bboxes,
#     plot_image,
#     save_checkpoint,
#     load_checkpoint,
# )
# from loss import YoloLoss
</span>
<span class="n">seed</span> <span class="o">=</span> <span class="mi">2023</span>
<span class="n">torch</span><span class="p">.</span><span class="n">manual_seed</span><span class="p">(</span><span class="n">seed</span><span class="p">)</span>

<span class="c1"># Hyperparameters etc.
</span><span class="n">LEARNING_RATE</span> <span class="o">=</span> <span class="mf">2e-5</span>
<span class="n">DEVICE</span> <span class="o">=</span> <span class="s">"cuda"</span> <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">cuda</span><span class="p">.</span><span class="n">is_available</span> <span class="k">else</span> <span class="s">"cpu"</span>
<span class="n">BATCH_SIZE</span> <span class="o">=</span> <span class="mi">16</span>
<span class="n">WEIGHT_DECAY</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">EPOCHS</span> <span class="o">=</span> <span class="mi">50</span>
<span class="n">NUM_WORKERS</span> <span class="o">=</span> <span class="mi">2</span>
<span class="n">PIN_MEMORY</span> <span class="o">=</span> <span class="bp">True</span>
<span class="n">LOAD_MODEL</span> <span class="o">=</span> <span class="bp">False</span>
<span class="n">LOAD_MODEL_FILE</span> <span class="o">=</span> <span class="s">"yolo_test.pth.tar"</span>


<span class="k">def</span> <span class="nf">train_fn</span><span class="p">(</span><span class="n">train_loader</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">loss_fn</span><span class="p">):</span>
    <span class="n">loop</span> <span class="o">=</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">train_loader</span><span class="p">,</span> <span class="n">leave</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="n">mean_loss</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="k">for</span> <span class="n">batch_idx</span><span class="p">,</span> <span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">loop</span><span class="p">):</span>
        <span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">),</span> <span class="n">y</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">loss</span> <span class="o">=</span> <span class="n">loss_fn</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
        <span class="n">mean_loss</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">loss</span><span class="p">.</span><span class="n">item</span><span class="p">())</span>
        <span class="n">optimizer</span><span class="p">.</span><span class="n">zero_grad</span><span class="p">()</span>
        <span class="n">loss</span><span class="p">.</span><span class="n">backward</span><span class="p">()</span>
        <span class="n">optimizer</span><span class="p">.</span><span class="n">step</span><span class="p">()</span>

        <span class="c1"># update progress bar
</span>        <span class="n">loop</span><span class="p">.</span><span class="n">set_postfix</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">item</span><span class="p">())</span>

    <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Mean loss was </span><span class="si">{</span><span class="nb">sum</span><span class="p">(</span><span class="n">mean_loss</span><span class="p">)</span><span class="o">/</span><span class="nb">len</span><span class="p">(</span><span class="n">mean_loss</span><span class="p">)</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">main</span><span class="p">():</span>
    <span class="n">model</span> <span class="o">=</span> <span class="n">YOLOv1</span><span class="p">(</span><span class="n">num_split_cell</span><span class="o">=</span><span class="mi">7</span><span class="p">,</span> <span class="n">num_boxes</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">num_classes</span><span class="o">=</span><span class="mi">3</span><span class="p">).</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">)</span>
    <span class="n">optimizer</span> <span class="o">=</span> <span class="n">optim</span><span class="p">.</span><span class="n">Adam</span><span class="p">(</span>
        <span class="n">model</span><span class="p">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="n">LEARNING_RATE</span><span class="p">,</span> <span class="n">weight_decay</span><span class="o">=</span><span class="n">WEIGHT_DECAY</span>
    <span class="p">)</span>
    <span class="n">loss_fn</span> <span class="o">=</span> <span class="n">YoloLoss</span><span class="p">()</span>

    <span class="k">if</span> <span class="n">LOAD_MODEL</span><span class="p">:</span>
        <span class="n">load_checkpoint</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">load</span><span class="p">(</span><span class="n">LOAD_MODEL_FILE</span><span class="p">),</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">)</span>

    <span class="n">train_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span>
        <span class="n">data_path</span><span class="o">=</span><span class="n">TR_DATA_PATH</span><span class="p">,</span>
        <span class="n">label_path</span><span class="o">=</span><span class="n">TR_LAB_PATH</span><span class="p">,</span>
        <span class="n">S</span> <span class="o">=</span> <span class="mi">7</span><span class="p">,</span>
        <span class="n">B</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span>
        <span class="n">C</span> <span class="o">=</span> <span class="mi">3</span><span class="p">,</span>
        <span class="n">transforms</span><span class="o">=</span><span class="n">get_transform</span><span class="p">()</span>
    <span class="p">)</span>

    <span class="n">val_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span>
        <span class="n">data_path</span><span class="o">=</span><span class="n">VAL_DATA_PATH</span><span class="p">,</span>
        <span class="n">label_path</span><span class="o">=</span><span class="n">VAL_LAB_PATH</span><span class="p">,</span>
        <span class="n">S</span> <span class="o">=</span> <span class="mi">7</span><span class="p">,</span>
        <span class="n">B</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span>
        <span class="n">C</span> <span class="o">=</span> <span class="mi">3</span>
    <span class="p">)</span>

    <span class="n">train_loader</span> <span class="o">=</span> <span class="n">DataLoader</span><span class="p">(</span>
        <span class="n">dataset</span><span class="o">=</span><span class="n">train_dataset</span><span class="p">,</span>
        <span class="n">batch_size</span><span class="o">=</span><span class="n">BATCH_SIZE</span><span class="p">,</span>
        <span class="n">num_workers</span><span class="o">=</span><span class="n">NUM_WORKERS</span><span class="p">,</span>
        <span class="n">pin_memory</span><span class="o">=</span><span class="n">PIN_MEMORY</span><span class="p">,</span>
        <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
        <span class="n">drop_last</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="n">val_loader</span> <span class="o">=</span> <span class="n">DataLoader</span><span class="p">(</span>
        <span class="n">dataset</span><span class="o">=</span><span class="n">val_dataset</span><span class="p">,</span>
        <span class="n">batch_size</span><span class="o">=</span><span class="n">BATCH_SIZE</span><span class="p">,</span>
        <span class="n">num_workers</span><span class="o">=</span><span class="n">NUM_WORKERS</span><span class="p">,</span>
        <span class="n">pin_memory</span><span class="o">=</span><span class="n">PIN_MEMORY</span><span class="p">,</span>
        <span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
        <span class="n">drop_last</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
    <span class="p">)</span>

    <span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">EPOCHS</span><span class="p">):</span>
        <span class="c1"># for x, y in train_loader:
</span>        <span class="c1">#    x = x.to(DEVICE)
</span>        <span class="c1">#    for idx in range(8):
</span>        <span class="c1">#        bboxes = cellboxes_to_boxes(model(x))
</span>        <span class="c1">#        bboxes = non_max_suppression(bboxes[idx], iou_threshold=0.5, threshold=0.4, box_format="midpoint")
</span>        <span class="c1">#        plot_image(x[idx].permute(1,2,0).to("cpu"), bboxes)
</span>
        <span class="c1">#    import sys
</span>        <span class="c1">#    sys.exit()
</span>
        <span class="n">pred_boxes</span><span class="p">,</span> <span class="n">target_boxes</span> <span class="o">=</span> <span class="n">get_bboxes</span><span class="p">(</span>
            <span class="n">train_loader</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">threshold</span><span class="o">=</span><span class="mf">0.4</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">DEVICE</span>
        <span class="p">)</span>

        <span class="n">mean_avg_prec</span> <span class="o">=</span> <span class="n">mean_average_precision</span><span class="p">(</span>
            <span class="n">pred_boxes</span><span class="p">,</span> <span class="n">target_boxes</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">box_format</span><span class="o">=</span><span class="s">"midpoint"</span>
        <span class="p">)</span>
        <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Train mAP: </span><span class="si">{</span><span class="n">mean_avg_prec</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>

        <span class="c1">#if mean_avg_prec &gt; 0.9:
</span>        <span class="c1">#    checkpoint = {
</span>        <span class="c1">#        "state_dict": model.state_dict(),
</span>        <span class="c1">#        "optimizer": optimizer.state_dict(),
</span>        <span class="c1">#    }
</span>        <span class="c1">#    save_checkpoint(checkpoint, filename=LOAD_MODEL_FILE)
</span>        <span class="c1">#    import time
</span>        <span class="c1">#    time.sleep(10)
</span>
        <span class="n">train_fn</span><span class="p">(</span><span class="n">train_loader</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">loss_fn</span><span class="p">)</span>
        <span class="n">torch</span><span class="p">.</span><span class="n">save</span><span class="p">(</span><span class="n">model</span><span class="p">.</span><span class="n">state_dict</span><span class="p">(),</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">WEIGHTS_PATH</span><span class="si">}</span><span class="s">yolov1_</span><span class="si">{</span><span class="n">EPOCHS</span><span class="si">}</span><span class="s">.pt'</span><span class="p">)</span>

<span class="k">if</span> <span class="n">__name__</span> <span class="o">==</span> <span class="s">"__main__"</span><span class="p">:</span>
    <span class="n">main</span><span class="p">()</span>
</code></pre></div></div>

<h2 id="6-evaluation">6. Evaluation</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">test_dataset</span> <span class="o">=</span> <span class="n">SoccerDataset</span><span class="p">(</span>
        <span class="n">data_path</span><span class="o">=</span><span class="n">TEST_DATA_PATH</span><span class="p">,</span>
        <span class="n">label_path</span><span class="o">=</span><span class="n">TEST_LAB_PATH</span><span class="p">,</span>
        <span class="n">S</span> <span class="o">=</span> <span class="mi">7</span><span class="p">,</span>
        <span class="n">B</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span>
        <span class="n">C</span> <span class="o">=</span> <span class="mi">3</span><span class="p">,</span>
        <span class="n">transforms</span> <span class="o">=</span> <span class="n">get_transform</span><span class="p">()</span>
    <span class="p">)</span>

<span class="n">test_loader</span> <span class="o">=</span> <span class="n">DataLoader</span><span class="p">(</span>
    <span class="n">dataset</span><span class="o">=</span><span class="n">test_dataset</span><span class="p">,</span>
    <span class="n">batch_size</span><span class="o">=</span><span class="n">BATCH_SIZE</span><span class="p">,</span>
    <span class="n">num_workers</span><span class="o">=</span><span class="n">NUM_WORKERS</span><span class="p">,</span>
    <span class="n">pin_memory</span><span class="o">=</span><span class="n">PIN_MEMORY</span><span class="p">,</span>
    <span class="n">shuffle</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
    <span class="n">drop_last</span><span class="o">=</span><span class="bp">True</span>
<span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">model</span> <span class="o">=</span> <span class="n">YOLOv1</span><span class="p">(</span><span class="n">num_split_cell</span><span class="o">=</span><span class="mi">7</span><span class="p">,</span> <span class="n">num_boxes</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">num_classes</span><span class="o">=</span><span class="mi">3</span><span class="p">).</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">load</span><span class="p">(</span><span class="sa">f</span><span class="s">'</span><span class="si">{</span><span class="n">WEIGHTS_PATH</span><span class="si">}</span><span class="s">yolov1_</span><span class="si">{</span><span class="n">EPOCHS</span><span class="si">}</span><span class="s">.pt'</span><span class="p">))</span>

<span class="n">pred_boxes</span><span class="p">,</span> <span class="n">target_boxes</span> <span class="o">=</span> <span class="n">get_bboxes</span><span class="p">(</span>
            <span class="n">test_loader</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">threshold</span><span class="o">=</span><span class="mf">0.4</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">DEVICE</span>
        <span class="p">)</span>

<span class="n">mean_avg_prec</span> <span class="o">=</span> <span class="n">mean_average_precision</span><span class="p">(</span>
    <span class="n">pred_boxes</span><span class="p">,</span> <span class="n">target_boxes</span><span class="p">,</span> <span class="n">iou_threshold</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">box_format</span><span class="o">=</span><span class="s">"midpoint"</span>
<span class="p">)</span>
<span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">"Test mAP: </span><span class="si">{</span><span class="n">mean_avg_prec</span><span class="si">}</span><span class="s">"</span><span class="p">)</span> <span class="c1"># 0.34의 mAP 기록. 과적합
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">im</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">test_dataset</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
<span class="n">im</span><span class="p">,</span> <span class="n">t</span> <span class="o">=</span> <span class="n">im</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">t</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
<span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
    <span class="n">prediction</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">im</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">))</span>
<span class="n">prediction</span> <span class="o">=</span> <span class="n">prediction</span><span class="p">.</span><span class="n">reshape</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">7</span><span class="p">,</span> <span class="mi">7</span><span class="p">,</span> <span class="mi">3</span> <span class="o">+</span> <span class="mi">2</span> <span class="o">*</span> <span class="mi">5</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">target</span> <span class="o">=</span> <span class="n">t</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">DEVICE</span><span class="p">)</span>
        <span class="c1"># cell당 2개의 box가 그려지기때문에 2개의 iou
</span><span class="n">iou_b1</span> <span class="o">=</span> <span class="n">intersection_over_union</span><span class="p">(</span><span class="n">prediction</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">],</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">])</span>
<span class="n">iou_b2</span> <span class="o">=</span> <span class="n">intersection_over_union</span><span class="p">(</span><span class="n">prediction</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">13</span><span class="p">],</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">])</span> <span class="c1"># target은 동일함
</span><span class="n">ious</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">iou_b1</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">iou_b2</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">0</span><span class="p">)],</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>

<span class="c1"># best
</span><span class="n">iou_maxes</span><span class="p">,</span> <span class="n">bestbox</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">ious</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span> <span class="c1"># 두 박스 중 iou 최대값 리턴, 몇번째 박스의 iou가 최대인지 idx 리턴,
</span><span class="n">exists_box</span> <span class="o">=</span> <span class="n">target</span><span class="p">[...,</span> <span class="mi">3</span><span class="p">].</span><span class="n">unsqueeze</span><span class="p">(</span><span class="mi">3</span><span class="p">)</span>  <span class="c1"># 해당 박스에 객체가 존재하는지. 존재하면 1, 아니면 0. 논문에서 말한 1obj_i
</span>
<span class="c1"># 1. localization loss 구하기
# 두 박스 중 bestbox를 고르는 방법
</span><span class="n">box_predictions</span> <span class="o">=</span> <span class="n">exists_box</span> <span class="o">*</span> <span class="p">(</span>
    <span class="p">(</span>
        <span class="n">bestbox</span> <span class="o">*</span> <span class="n">prediction</span><span class="p">[...,</span> <span class="mi">9</span><span class="p">:</span><span class="mi">13</span><span class="p">]</span> <span class="c1"># best box가 1이라면 두번째 박스를 골라야하니까 9:13을
</span>        <span class="o">+</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="n">bestbox</span><span class="p">)</span> <span class="o">*</span> <span class="n">prediction</span><span class="p">[...,</span> <span class="mi">4</span><span class="p">:</span><span class="mi">8</span><span class="p">]</span> <span class="c1"># best box가 0이라면 첫번째 박스를 골라야하니까 4:8을 선택.
</span>    <span class="p">)</span>
<span class="p">)</span>

<span class="n">class_predictions</span> <span class="o">=</span> <span class="n">exists_box</span> <span class="o">*</span> <span class="n">prediction</span><span class="p">[...,:</span><span class="mi">3</span><span class="p">]</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">img</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">(</span><span class="n">im</span><span class="p">.</span><span class="n">squeeze</span><span class="p">(</span><span class="mi">0</span><span class="p">).</span><span class="n">permute</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">255</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">uint8</span><span class="p">).</span><span class="n">copy</span><span class="p">()</span>

<span class="k">for</span> <span class="n">y_idx</span><span class="p">,</span> <span class="n">y</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">box_predictions</span><span class="p">[</span><span class="mi">0</span><span class="p">]):</span>
  <span class="k">for</span> <span class="n">x_idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">y</span><span class="p">):</span>
    <span class="k">if</span> <span class="n">x</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span> <span class="o">!=</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># width가 0이 아니면 박스가 그려진 것
</span>        <span class="n">x_cell</span><span class="p">,</span> <span class="n">y_cell</span><span class="p">,</span> <span class="n">w_cell</span><span class="p">,</span> <span class="n">h_cell</span> <span class="o">=</span> <span class="n">box_predictions</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span><span class="n">y_idx</span><span class="p">,</span><span class="n">x_idx</span><span class="p">].</span><span class="n">tolist</span><span class="p">()</span>
        <span class="n">label</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">class_predictions</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span><span class="n">y_idx</span><span class="p">,</span><span class="n">x_idx</span><span class="p">]).</span><span class="n">tolist</span><span class="p">()</span>
        <span class="n">x1</span><span class="p">,</span> <span class="n">y1</span> <span class="o">=</span> <span class="p">(</span><span class="n">x_cell</span> <span class="o">+</span> <span class="n">x_idx</span><span class="p">)</span> <span class="o">/</span> <span class="mi">7</span><span class="p">,</span> <span class="p">(</span><span class="n">y_cell</span> <span class="o">+</span> <span class="n">y_idx</span><span class="p">)</span> <span class="o">/</span> <span class="mi">7</span> <span class="c1"># x_cell, y_cell 형식을 x_center, y_center 형식으로 변환
</span>        <span class="n">w</span><span class="p">,</span> <span class="n">h</span> <span class="o">=</span> <span class="n">w_cell</span> <span class="o">/</span> <span class="mi">7</span><span class="p">,</span> <span class="n">h_cell</span> <span class="o">/</span> <span class="mi">7</span>
        
        <span class="n">float_x_center</span> <span class="o">=</span> <span class="mi">448</span> <span class="o">*</span> <span class="n">x1</span>
        <span class="n">float_y_center</span> <span class="o">=</span> <span class="mi">448</span> <span class="o">*</span> <span class="n">y1</span>
        <span class="n">float_width</span> <span class="o">=</span> <span class="mi">448</span> <span class="o">*</span> <span class="n">w</span>
        <span class="n">float_height</span> <span class="o">=</span> <span class="mi">448</span> <span class="o">*</span> <span class="n">h</span>

        <span class="n">x1</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">float_x_center</span> <span class="o">-</span> <span class="n">float_width</span> <span class="o">/</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">y1</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">float_y_center</span> <span class="o">-</span> <span class="n">float_height</span> <span class="o">/</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">x2</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">float_width</span><span class="p">)</span> <span class="o">+</span> <span class="n">x1</span>
        <span class="n">y2</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">float_height</span><span class="p">)</span> <span class="o">+</span> <span class="n">y1</span>
        
        <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">img</span><span class="p">,</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span> <span class="n">y1</span><span class="p">),</span> <span class="p">(</span><span class="n">x2</span><span class="p">,</span><span class="n">y2</span><span class="p">),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
        <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">img</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">label</span><span class="p">),</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="o">+</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">img</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/69788b4c-3c80-420e-9b41-9d9875e22d33" />
  </p>
</div>

<p><br /><br /></p>

<p>최종적으로 학습에서는 0.85의 mAP, test에서는 0.34의 mAP를 기록하며 과적합된 모델이다. 아마 데이터셋의 크기가 작은 것이 과적합에 큰 영향을 미쳤을 것 같다. 위의 이미지를 보면 box prediction 부분에서 아쉬움을 보인다. 하지만, 학습 속도는 한 에포크당 20초가 걸렸고 이는 SSD와 비슷한 속도이면서 Faster R-CNN보다 약 10배이상 빠른 속도이다.</p>

<p>논문에서 언급한대로 Faster R-CNN에 비해 mAP 측면에서는 좀 아쉽긴하지만 속도면에서는 매우 빠른 장점을 가진 YOLOv1 모델이다.</p>

<h1 id="이미지-출처">이미지 출처</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1506.02640v5.pdf">You Only Look Once: Unified, Real-Time Object Detection Paper</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[Abstract 객체 탐지의 이전 연구들에서는 분류기에서도 detection을 수행할 수 있도록 했다. 본 연구에서는 하나의 단일 신경망을 사용해서 여러 bounding box와 class probabilities를 예측하는 객체 탐지를 regression 문제로 취급했다. 본 연구의 YOLO 모델은 객체 탐지 파이프라인 전체가 하나의 네트워크로 이루어져있어서 실시간 이미지를 초당 45프레임으로 처리할만큼 굉장히 빠르다. 다른 네트워크와 비교했을 때 localization에 대한 오류가 존재하긴하지만 배경에 대해서는 예측을 잘 수행하고 다른 도메인에 대해서도 RCNN, DPM보다 일반화가 잘 된 예측을 수행한다.]]></summary></entry><entry><title type="html">RCNN : Rich Feature Hierarchies for Accurate Object Detection and Semantic Segmentation 논문 요약</title><link href="https://kimyeonz.github.io/blog/detection/RCNN/" rel="alternate" type="text/html" title="RCNN : Rich Feature Hierarchies for Accurate Object Detection and Semantic Segmentation 논문 요약" /><published>2023-10-13T00:00:00+00:00</published><updated>2023-10-13T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/detection/RCNN</id><content type="html" xml:base="https://kimyeonz.github.io/blog/detection/RCNN/"><![CDATA[<h1 id="abstract">Abstract</h1>
<p>지난 몇년동안(2013) PASCAL VOC dataset에 관해 객체 탐지의 성능 향상이 더디다. 대화마다 best model은 일반적으로 여러개의 저수준 피처와 고수준의 context를 결합한 복잡한 앙상블 구조이다. 
본 연구에서는 단순한 탐지 알고리즘으로 VOC 2012에서 mAP 53.3%를 달성하며 지난 대회 대비 30% 이상 성능을 향상시킨 방법에 대해 제안한다. 본 논문의 두가지 키포인트는 첫째, bottom-up region proposal에 
high capacity CNN을 적용해서 localize와 segment objects를 할 수 있는 것과 둘째, 훈련 데이터가 부족한 경우 사전 학습 모델을 fine tuning하여 성능을 크게 높였다는 것이다. Region
proposal이 CNN과 결합되었기 때문에 R-CNN이라 부르고, 최근 개발된 sliding window 개념을 활용한 CNN 아키텍처인 OverFeat과 비교했을 때 ILSVRC2013 탐지부문에서 큰 격차로 우승했다.
<br /><br /></p>

<h1 id="details">Details</h1>
<h2 id="도입부">도입부</h2>
<ul>
  <li>지난 10년간 visual recogniton task에서 SIFT, HOG와 같은 방법을 기반으로 많은 진전이 있었다. 하지만, 2010 ~ 2012년까지 일반적인 객체 탐지 부문인 PASCAL VOC에서는 발전이 더디었다.</li>
  <li>CNN이 ImageNet 데이터뿐만 아니라 PASCAL VOC에도 적합할까?</li>
  <li>본 논문은 deep network로 객체를 localize하고, 적은 데이터만 가지고도 high-capacity model을 훈련하는 두 가지 문제에 집중하며 CNN이 HOG에 비해 PASCAL VOC의 객체 탐지 성능을 크게 향상시켰다는 것을 보여주는 첫번째 논문이다.</li>
  <li>이미지 분류와는 달리 객체 탐지에서는 localization이라는 회귀 문제가 있기때문에 기존의 CNN 방식으로는 성능이 그리 좋지 않았고, sliding window 개념을 사용하여 공간 정보를 잘 유지하려고했지만,
layer가 깊어질수록 입력이미지에 대한 receptive field가 커져서 정밀한 localization을 수행하기에는 어려움이 따른다.</li>
  <li>이를 해결하기 위해 약 2000개의 독립적인 RoI(Region of Interest, 물체가 존재할 것이라고 판단하는 영역)를 제시하고 각 영역에 대해 CNN으로 고정된 길이의 특징을 추출한 다음
추출된 결과를 SVM을 통해 클래스로 분류하는 방법을 사용한다.</li>
  <li>객체 탐지 시 데이터가 충분하지 않아서 사전학습 된 모델을 fine tuning하여 성능을 8%나 향상시켰다.</li>
</ul>

<p><br /><br /></p>

<h2 id="object-detection-with-r-cnn">Object detection with R-CNN</h2>
<p>본 연구의 object detection은 크게 3가지 모듈로 이루어져있다. 첫째는 region proposal을 생성하는 모듈로 객체가 있을 것이라고 판단되는 region의 후보를 추려내는 모듈이다. 둘째는,
CNN을 통해 추려낸 모듈 중에 고정된 길이의 특징 벡터를 추출하는 모듈이다. 셋째는, SVM을 통해 분류를 수행하는 모듈이다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/eddb0601-4be5-4d12-8c3f-467dd3d36d6e" />
  </p>
</div>

<p><br /><br /></p>

<h3 id="1-module-design">1. Module design</h3>
<p><strong>Region proposals</strong></p>
<ul>
  <li>최근 많은 논문들이 cateogy-independent region proposals를 생성하는 방법을 제공하고 있다.</li>
  <li>R-CNN은 selective search(이미지의 얼룩진 부분을 유사한 category로 판단하여 하나의 영역으로 제안하는 알고리즘. 그 후 bottom-up 방식으로 비슷한 영역을 합쳐서 결과적으로 약 2000개의 region proposal을 만듬)를 기반으로 region proposals을 수행한다.</li>
</ul>

<p><strong>Featrue extraction</strong></p>
<ul>
  <li>AlexNet 연구를 기반으로 selective search로 제안한 영역에서 4096의 차원을 갖는 특징을 추출한다. 5개의 합성곱과 2개의 전결합층으로 이루어져있다.</li>
  <li>제안된 영역은 CNN의 구성에 맞게 크기가 조정되어야 하므로(227 x 227) bonding box의 크기나 가로 세로 비율에 상관없이 warping(crop, 왜곡 등의 이미지 변형)시킨다.</li>
  <li>경계에 bounding box가 있으면 약간 확대시키고 padding 16을 적용한다.</li>
  <li>Appendix A
    <ul>
      <li>CNN에 region proposals를 입력하기 전 CNN input에 적합한 형태로 변환이 필요하고 총 3가지 방법을 사용했다. 아래 사진 (A)는 original region proposals이다.</li>
      <li>첫번째 방법 (B)는 이미지 사이즈를 줄여서 특정 크기의 정사각형 안에 이미지를 입력하되 객체 외의 주변 context를 조금 포함시키는 방법이다.</li>
      <li>두번째 방법 (C)는 정사각형 안에 이미지를 입력하되 객체 외의 주변 context를 포함시키지 않는 방법이다.</li>
      <li>세번째 방법 (D)는 정사각형 안에 이미지를 입력하되 가로세로 비율을 고려하지 않고 이미지를 warped 시키는 방법이다.</li>
      <li>각 이미지의 위의 행은 padding=0, 아래 행은 padding=16을 사용한 경우이다.</li>
      <li>만약 주어진 정사각형이 region proposals보다 크다면, missing data는 image mean으로 대체했다. 대체 된 부분은 CNN 입력전 다시 제거한다.</li>
    </ul>
  </li>
</ul>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/7e8f375b-bf63-4cb4-97d9-c0de3fc744f4" />
  </p>
</div>

<p><br /><br /></p>

<h3 id="2-test-time-detection">2. Test-time detection</h3>
<p>Test 시에는 제안된 영역마다 특징이 추출되면 해당 특징을 SVM에 입력해 각 영역마다의 점수를 매긴다. 최종 예측을 위해서는 하나의 객체당 하나의 bonding box만 남겨져야하는데 NMS(Non-Maximum Suppression) 방법을 통해 confidence score가 threshold보다 낮은 박스를 모두 제거시키고, 만약 threshold보다 높더라도 confidence score가 높은 bounbding box와 많이 중첩된(IoU가 높은) box(condifence score가 덜 높은)가 있다면 confidence score가 가장 높은 box를 제외하고 모든 box를 제거시켜 각각의 객체마다 하나의 bounding box만 남도록 한다.</p>

<p><strong>NMS(Non-Maximum Supression)</strong>
<br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/1cfd20dc-e438-498e-b40d-2850808ee9ee" />
  </p>
</div>

<p><br /></p>

<p>R-CNN에서는 selective search 알고리즘을 통해 region proposals을 하고 이 region들을 CNN에 넣어 특징을 추출한다. 추출된 특징은 SVM과 box linear regression을 통해 region 내의 box에 존재하는 객체가 특정 categories에 속할 확률을 예측하고 box의 위치를 ground truth box와 가깝도록 조정한다. 하지만 box에 대해 클래스를 잘 예측하고 box를 잘 그렸다고 할지라도 위의 그림처럼 같은 객체를 여러개의 box가 예측하고 있을 수도 있다. 이는 연산량의 측면에서도 매우 비효율적이기떄문에 가장 객체를 잘 담은 box만 남겨두고 다른 box를 모두 제거하는게 좋을 것이다.</p>

<p>NMS는 가장 좋은 box만을 남겨두기 위한 기법으로 다음과 같은 과정으로 진행된다.</p>

<p>0) 각각의 box마다 계산된 confidence score를 기준으로 confidence score가 threshold를 넘지 않는다면 모두 제거한다. 여기서 말한 confidence score는 물체 존재 확신도로 일반적으로는 box내 객체가 속할 점수들 중 가장 큰 점수로 선택된다(개 0.5, 고양이 0.3, 사자 0.2로 softmax의 예측이 수행되었다면 개의 확률인 0.5가 confidence score로 사용됨). Confidence score를 통한 box 제거는 NMS와 분리된 단계로도 볼 수 있기때문에 0 단계로 표현했다.</p>

<p>1) 모든 box를 confidence score를 기준으로 내림차순 정렬한다.</p>

<p>2) 정렬된 box중 가장 위에 있는 box를 선택하고 이 box와 다른 box와의 IOU를 계산한다.</p>

<p>3) IoU가 threshold보다 높다면 즉, 가장 confidence score가 높은 box와 너무 많이 겹친 박스라면 같은 객체를 담고 있을 확률이 크기때문에 제거한다.</p>

<p>4) 이 과정을 반복하며 객체마다 하나의 box만 남게된다.</p>

<p>Confidence score threshold가 높다면 물체가 존재한다는 기준을 까다롭게 잡은 것이니까 0단계에서 많은 box들이 제거될 것이고, IoU threshold가 낮다면 조금만 겹쳐도 같은 객체를 예측한다고 판단하여 많은 box가 제거될 것이다. 따라서, 적절한 threshold를 선택하는 것도 중요하다.</p>

<p>NMS 이후에는 이미지내의 다른 object들을 기반으로 최종 박스들을 rescoring 하여 결과를 출력한다.</p>

<p><strong>Run-time analysis</strong></p>
<ul>
  <li>모든 CNN 파라미터가 모든 categories에 대해 공유되고, 특징 벡터가 다른 접근 방식들에 비해 저차원이라는 2가지 특징이 detection을 효과적으로 만들었다.</li>
  <li>Class별로 이루어지는 유일한 계산은 특징 벡터와 SVM 가중치 사이의 연산과 NMS이다. Feature matrix는 일반적으로 2000 x 4096이고, SVM 가중치 행렬은 4096 x N(N은 클래스수)이다.</li>
  <li>고차원의 특징을 학습할 때보다 훨씬 더 적은 연산시간을 가진 효과적인 네트워크를 구축했다. 더 좋은 성능을 가지면서 다른 모델에 비해 예측 시간이 빠른 모델이다.</li>
</ul>

<h3 id="3-training">3. Training</h3>
<p><strong>Supervised pre-training &amp; Domain-specific fine-tuning</strong></p>
<ul>
  <li>Bounding box는 없지만 ILSVRC 2012 classification datsets을 가지고 CNN 모델을 pre-training 시켰다.</li>
  <li>Detection이라는 새로운 tasks를 수행하기위해 SGD를 사용해 CNN 파라미터를 fine-tuning 시켰고 데이터는 warped region proposals만 사용했다.</li>
  <li>IoU의 threshold는 bonding box의 IoU가 0.5이상을 positive, 미만을 negative로 했고 pre-trained 모델의 1/10인 0.001을 SGD의 학습률로 사용했다. 각 SGD iteration에는
128(32개의 전경 class, 96개의 배경 class)개의 mini batch가 사용되었다.</li>
</ul>

<p><strong>Object category classifier</strong></p>
<ul>
  <li>자동차를 탐지하는 task라면 자동차를 잘 담고 있는 box를 positive, 그 외 배경을 담고 있으면 negative로 취급한다. 하지만, 만약 자동차를 일부만 포함하고 있는 box가 있다면 positive라고 할 수 있을까?</li>
  <li>본 연구에서는 이와 같은 overlap 문제를 해결하기 위해 regions 중 IoU overlap threshold가 0.3이상이면 해당 box에 객체가 있다고 판단한다(자동차의 0.3% 정도 담고있으면 해당 부분을 positive로 인정). 0.3이라는 수치는 validation set에서 grid search(0.1, …, 0.5)를 통해 도출했다.</li>
  <li>특징이 추출되면 각 클래스마다 하나의 linear SVM을 사용해서 정답과 예측결과를 비교하며 모델을 최적화시키는데 모든 클래스를 한번에 처리하지 않고, 각각의 클래스마다 독립적으로 SVM을 최적화시키기때문에 많은 시간이 소요된다.</li>
  <li>이를 해결하기위해 수렴이 빠른 standard hard negative mining methods(negative가 더 많은 불균형 데이터에는 FP오류가 많을 것이기때문에 모델이 FP로 잘못 예측한 데이터를 select하여 학습에 다시 사용해서 모델을 좀 더 강건하게 만드는 방법)를 사용했다.</li>
  <li>Appendix B
    <ul>
      <li>CNN에 softmax layer를 추가하여 class 예측을 수행하는 fine-tuning 방법(IoU overlap threshold 0.5)과 기존처럼 SVM(IoU overlap threshold 0.3)을 학습할 때 positive, negative에 대한 기준이 달라진 이유는 fine-tuning과 SVM에서 동일한 IoU overlap threshold 사용했을 때 성능이 오히려 하락했는데 아무래도 fine-tuning에는 과대적합을 방지하기 위해 더 많은 데이터들이 필요하지만 현실적으로 데이터를 추가시키기 어렵기때문에 threshold를 SVM과 통일시키지 못했다.</li>
      <li>Fine-tuning에서의 분류기를 쓰지 않고 SVM을 굳이 학습시킨 이유는 fine-tuning에서 positive class에 대해 localization이 정확하게 수행되지 않았고, softmax classifier가 SVM처럼의 hard negative mining의 범위가 아닌 전체 범위에서 무작위로 negative class를 추출했기때문이라고 추측된다.</li>
    </ul>
  </li>
</ul>

<h3 id="4-results-on-pascal-voc-2010-12--ilsvrc2013-detection">4. Results on PASCAL VOC 2010-12 &amp; ILSVRC2013 detection</h3>
<ul>
  <li>다른 여러 모델들과 비교했을 때 본 연구의 SVM은 VOC 2010에서 53.7% mAP를, 2011/12에서 53.3% mAP를 달성하며 가장 좋은 성능을 냈다.</li>
  <li>ILSVRC2013 detection에서는 OverFeat의 24.3%의 mAP보다 훨씬 더 뛰어난 31.4%의 mAP를 기록하며 이전 대회의 기록보다 더 좋은 성능을 냈다.</li>
</ul>

<p><br /><br /></p>

<h2 id="visualization-ablation-and-modes-of-error">Visualization, ablation, and modes of error</h2>
<h3 id="1-visualizaing-learned-features">1. Visualizaing learned features</h3>

<p><br /></p>
<div align="center">
  <p>
  <img width="800" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/164746a1-7acb-462c-8182-8c50e2ac25e7" />
  </p>
</div>

<p><br /></p>

<ul>
  <li>첫번째 레이어는 선, 점과 같이 저수준의 특징을 학습하기때문에 직관적인 시각화가 가능한데 그 후의 레이어들은 복잡한 특징을 학습해서 시각화하기가 어렵다.</li>
  <li>따라서, 각층마다 어떻게 학습이 이루어지는지 파악하기 위해 non-parametric(비모수적 방법. 분석에 대해 사전가정을 포함하지 않는, 즉, 파라미터가 사전에 정해져있지 않은 통계적 기법. 데이터에서 패턴이나
관계를 추정할 때 사용되고 데이터의 순위, 순서, 통계량 등 데이터를 기반으로 분석을 진행)방법을 사용한다.</li>
  <li>특정 feature(unit)을 하나의 detector로 취급하고 여러 region proposals에 대해 activation을 적용해서 activation이 높은 순부터 낮은 순까지 정렬한 후 NMS를 적용하고 top score를 시각화한다.
이 방법은 activation이 높은, 즉, top ranking에 위치한 unit은 내가 물체를 가지고 있다는 자신감을 대변한다는 idea다.</li>
  <li>layer pool5의 units을 시각화한 결과, network는 몇몇 클래스의 feature와, 모양, 텍스처, 색상 등을 결합하여 학습한다는 것을 알 수 있다. 이후, fully connected layer는 이 정보들을 기반으로 더 높은 수준의 특징을 가지고 학습을 진행한다.</li>
</ul>

<h3 id="2-ablation-studies">2. Ablation studies</h3>
<p><strong>Performance layer-by-layer, without fine-tuning</strong></p>
<ul>
  <li>어떤 레이어가 중요한지 파악하기 위해 CNN의 마지막 3개 레이어를 살펴봤다(Layer pool5는 생략).</li>
  <li>Layer fc6은 4096 x 9216에 pool5의 피처맵을 곱하고, Layer fc7은 4096 x 4096에 fc6의 피처맵을 곱한다.</li>
  <li>fine tuning을 사용하지 않았을 때는 Layer fc6과 fc7이 없을 때 성능이 더 좋았다. 전결합층으로 인한 많은 연산량이 오히려 성능을 하락시켰다.</li>
  <li>CNN은 전결합층보다 convolution layer가 중요하다.</li>
</ul>

<p><strong>Performance layer-by-layer, with fine-tuning</strong></p>
<ul>
  <li>Fine tuning 결과 성능이 8%나 상승한걸로보아 fine tuning의 효과가 pool5보다 fc6,7에서 더 컸다고 할 수 있다.</li>
  <li>이로인해 pool5까지는 ImageNet을 학습하여 일반적인 특징이 추출되고, fine tuning한 분류기에서 해당 도메인에 특화되고 non-linear한 분류기가 구축된 것을 알 수 있다.</li>
</ul>

<p><strong>Comparison to recent feature learning methods</strong></p>
<ul>
  <li>DPM에서 사용한 DPM ST와 DPM HSC 이 두 가지 feature learning methods를 RCNN에 적용하여 DPM과 결과를 비교했다.</li>
  <li>Feature learning methods를 사용한 RCNN은 최신 DPM과 비교했을 때보다 더 좋은 성능을 가졌다.</li>
</ul>

<h3 id="3-network-architectures">3. Network architectures</h3>
<ul>
  <li>본 연구에서 구현된 대부분의 아키텍처는 AlexNet을 참고했지만 어떤 아키텍처를 쓰느냐에 따라 탐지 성능이 크게 달라진다는 것을 알게되었다.</li>
  <li>O-Net(OxfordNet, VGG16)을 pre-trained 모델로 사용하고, 같은 환경에서 pre-trained 시킨 T-Net(TorontoNet)과 비교해본 결과 O-Net의 성능이 더 뛰어났다.</li>
  <li>하지만, O-Net은 T-Net에 비해 7배 더 큰 연산 시간을 가진다는 것이 한계로 드러났다.</li>
</ul>

<h3 id="4-detection-error-analysis--bonding-box-regression">4. Detection error analysis &amp; Bonding-box regression</h3>
<ul>
  <li>Hoiem의 Detection analysis를 기반으로 본 모델의 에러 모드와, fine-tuning이 에러 모드를 어떻게 바꾸는지, 에러 type이 DPM과 어떻게 다른지 비교했다.</li>
  <li>Error anlysis를 기반으로 localization error를 줄이는 방법을 적용했다. DPM의 bounding-box regression에 영감을 받아, selective search region proposal에 대한
pool5의 features가 주어지면 새로운 detection window를 예측하는 linear regression 모델을 학습시킨다.</li>
  <li>이 간단한 접근 방식이 mislocaliztion을 줄여서 mAP를 3~4 points 올리는 중요한 방법이 되었다.</li>
  <li>Appendix C
    <ul>
      <li>SVM으로 class 분류가 끝나면 class-specific bounding box regressor로 예측을 수행한다.</li>
      <li>$P$는 pixel, $G$는 ground-truth이다.</li>
      <li>$P_x, P_y$에 대해서는 점이기때문에 scale은 그대로 한채 변환하고, $P_w,P_h$에 대해서는 log 변환을 수행한다.</li>
      <li>
        <p>위의 $P$들을 얼만큼 이동시킬 것인지에 대한 정보는 $d(P)$를 통해 구할 수 있고, 이 $d(P)$ 정보를 기반으로 ground-truth에 가까운 예측을 수행하게된다. 식은 다음과 같이 표현할 수 있다.</p>

        <p><br /></p>
        <div align="center">
  <p>
  <img width="300" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/b201a28b-30cd-4ca0-81b0-1ae4c99e7ef0" />
  </p>
</div>

        <p><br /></p>
      </li>
      <li>$d(P)$들은 pool5 features들을 이용한 linear function이고 여기서 pool5의 features는 $\phi_5(P)$로 표현한다.</li>
      <li>linear function인 $d(P) = w^T \phi_5(P)$이고 $w$는 학습가능한 파라미터이다. $w$ 파라미터를 최적화시키기 위해서는 다음과 같이 정의한다.</li>
    </ul>

    <p><br /></p>
    <div align="center">
    <p>
    <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/c59240af-2388-4138-9145-3e43a421b4b8" />
    </p>
  </div>

    <p><br /></p>

    <ul>
      <li>위의 식에서 $t$는 ground-truth G와 예측 P의 차이로 정답에 가까이가기 위해서 P를 얼마나 조정해야하는지 알려준다. $w$의 식을 보면 결국 선형회귀인 MSE와 유사한 손실함수이고, 과적합을 방지하기 위해 regularization parameter인 lambda를 적용했다.</li>
    </ul>

    <p><br /></p>
    <div align="center">
    <p>
    <img width="400" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/ed053334-aea1-4387-8b2c-22a84a793511" />
    </p>
  </div>

    <p><br /></p>

    <ul>
      <li>쉽게 말하면, CNN을 통해 추출된 피처와 각 point마다 얼만큼 이동시켜야하는지에 대한 정보인 가중치 $w$를 곱해서 bounding box를 업데이트시키는 방법으로 선형 회귀를 학습하는 것이다.</li>
      <li>Bounding-box regression을 통해 regularization이 매우 중요하다는 것을 알았고 labmda = 1000으로 설정했다.</li>
      <li>Region Proposal이 ground-truth와 너무 많이 떨어져있다면 예측이 매우 힘들어지기때문에 IoU overlap threshold를 0.6으로 설정해서 많이 겹치지 않은 proposal은 모두 제거했다.</li>
    </ul>
  </li>
</ul>

<p><br /><br /></p>

<h2 id="the-ilsvrc2013-detection-dataset">The ILSVRC2013 detection dataset</h2>
<h3 id="1-dataset-overview">1. Dataset overview</h3>
<ul>
  <li>학습 약 39만개/검증 약 2만개/테스트 약 4만개</li>
  <li>검증과 테스트 데이터에는 라벨링이 잘 되어있지만 학습 데이터에는 모든 라벨들에 대해 라벨링이 잘 되어있진 않다.</li>
  <li>주로 검증 데이터에 의존하고, 학습 데이터는 보조로 사용한다. 검증 데이터를 동일한 크기를 갖게 val1(학습에 사용), val2(검증에 사용)로 나누어서 학습과 검증에서 사용한다.</li>
  <li>클래스가 불균형하기때문에 검증 데이터셋을 분리할 때 최대한 밸런스하게 분할했다.</li>
</ul>

<h3 id="2-region-proposals">2. Region proposals</h3>
<ul>
  <li>PASCAL과 같은 방법으로 region proposals를 수행했다.</li>
  <li>학습 데이터를 제외한 val1, val2, test 데이터에 selective search를 적용했다.</li>
  <li>ILSVRC image의 사이즈 범위는 매우 넓어서 selective search 전에 이미지 넓이를 500 pixels로 조정했다.</li>
  <li>검증 데이터에 selective search를 적용한 결과 이미지당 평균 2403개의 region propsals이 나왔고 이는 91.6%의 recall을 기록했다.</li>
  <li>PASCAL 데이터에서 98%의 recall에 비해 현저히 낮은 수치이기때문에 이로인해 region proposals 단계가 매우 중요하다는 것을 알 수 있다.</li>
</ul>

<h3 id="3-training-data">3. Training data</h3>
<ul>
  <li>학습 데이터를 구축할 때는 val1과 각 클래스당 N개의 ground-truth boxes만큼의 데이터를 추출했고, 만약 특정 클래스가 minor 클래스라서 N개보다 적다면 해당 클래스의 데이터는 모두 선택했다. $val1 + train_N$ 의 set으로 구성된다.</li>
  <li>학습 데이터는 CNN fine-tuning, detector SVM training, bounding-box regressor training이라는 3가지 task에 사용된다.</li>
  <li>Hard negative mining은 val1으로부터 random하게 5000개의 샘플을 뽑아 진행했다.</li>
</ul>

<h3 id="4-validation-and-evaluation--ablation-study">4. Validation and evaluation &amp; Ablation study</h3>
<ul>
  <li>결과를 제출하기 전에 val1+train, val2와 같은 data usage와 fine tuning, bounding box regression의 효과를 검증했다.</li>
  <li>일반화 능력을 판단하기위해 PASCAL과 똑같은 하이퍼파라미터(NMS threshold, SVM C, padding 등)로 실험했다.</li>
  <li>Bounding box regression이 있는 버전과 없는 버전으로 submission 했다.</li>
  <li>Fine tuning, val1 데이터셋에 train 데이터 추가, bounding box regression을 도입할 때 위 3가지 중 아무런 방법을 사용하지 않은 모델보다 성능이 더 좋았다.</li>
</ul>

<h3 id="5-relationship-to-overfeat">5. Relationship to OverFeat</h3>
<ul>
  <li>RCNN과 OverFeat은 구조적으로 굉장히 유사하지만 CNN을 이용한 fine tuning, SVM을 사용했다는 점에서 차이가 존재한다.</li>
  <li>OverFeat은 속도가 RCNN보다 9배나 빠른데 이 속도는 sliding windows(ex. region proposal)를 진행할 때 이미지가 warp 되지않아서 연산이 훨씬 쉽기 때문이다.</li>
  <li>RCNN도 다양한 방법을 동원해서 속도 문제를 개선해야한다.</li>
</ul>

<p><br /><br /></p>
<h2 id="semantic-segmentation">Semantic segmentation</h2>
<p>본 연구 당시 Semantic segmentation 분야에서 가장 좋은 모델이라고 평가받는 second order pooling 기법을 활용한 O2P 모델과 성능을 비교하기 위해 그들이 사용한 오픈 소스 프레임워크를 사용했다. O2P는 CPMC 기술을 사용하여 이미지 당 150개의 region proposals을 생성하고 SVR(support vector regression)을 사용하여 localization을 수행했다.</p>

<p><strong>CNN features for segmentation</strong></p>
<ul>
  <li>CPMC 알고리즘으로 추출한 regions의 features을 계산하기 위해 3가지 방법을 사용했다.</li>
  <li>첫째, region의 형태를 무시하고 warped window에 CNN features를 바로 계산하는 full 전략이다.</li>
  <li>둘째, foregound mask에 대해서만 CNN에 넣어 features를 계산하는 fg(foreground) 전략이다. 배경 region은 zero로 만드는 정규화 기법을 사용한다.</li>
  <li>셋째, 두 가지 방법을 모두 섞은 full+fg 전략이다.</li>
</ul>

<p><strong>Results on VOC 2011</strong></p>
<ul>
  <li>모든 feature computation strategy에서 fc6까지 사용하는게 fc7보다 성능이 좋았다.</li>
  <li>full+fg인 3번째 전략을 사용했을 때 성능이 가장 좋았고, 21개 중 11개의 카테고리에 대한 예측력이 매우 좋았다.</li>
  <li>RCNN은 R&amp;P(Region &amp; Parts)와 O2P 모델보다 성능이 뛰어났으며, fine tuning을 통해 추가로 성능이 향상될 것으로 기대한다.</li>
</ul>

<p><br /><br /></p>

<h1 id="conclusion">Conclusion</h1>
<ul>
  <li>최근 성능 향상이 더디던 객체 탐지 분야에서 기존보다 30% 개선된 성능을 낸 모델을 개발했다.</li>
  <li>Bottom-up region proposals을 CNN에 적용하여 localization과 segments를 수행했고 labeled data가 부족해서 image classification에 사용한 모델을 fine tuning하여 detection 성능을 향상시켰다.</li>
  <li>Supervised pre-training과 domain specific fine tuning 패러다임이 데이터가 부족한 문제에서 굉장히 큰 효과를 냈다고 추측한다.</li>
  <li>CV분야의 툴과 딥러닝을 결합(Bottom-up region proposals과 CNN의 결합)하는 방법론은 본 연구에 큰 도움이 됐다.</li>
</ul>

<p><br /><br /></p>

<h1 id="개인적인-생각">개인적인 생각</h1>
<ul>
  <li>Detection은 classification에 비해 더 복잡한 task인데 데이터 부족 문제를 해결하기위해 image classification model을 사전훈련 모델로 사용했고, fine tuning 결과 성능을 크게 향상시킨 것을 보아 딥러닝이 발전할수록 맨땅에 헤딩하는 심정으로 밑바닥부터 모델을 구현한다기보다 이전에 개발된 모델들을 어떻게 효과적으로 사용하는지가 더 효율적인 시대가 된 것 같다.</li>
  <li>Features를 처리하는 방법, layer의 깊이, dataset의 크기 등 하나의 연구에 가장 효과적인 case가 무엇인지 수많은 실험을 했고, 이를 통해 본 연구에서 객체 탐지 부문의 성능을 높이기위해 얼마나 많은 시간과 노력을 쏟았는지 알 수 있었다.</li>
  <li>하나의 이미지를 처리하는데 약 50초라는 긴 시간이 걸렸는데 향후 개발된 Fast-RCNN, Faster-RCNN에서는 어떻게 이 시간을 크게 단축시켰는지 기대가 된다.</li>
  <li>본 연구팀도 각 class마다 SVM 모델을 따로 구축하는 것은 매우 비효율적이라고 생각했을것 같은데 hard negative mining을 포기하고 딥러닝 네트워크에서 softmax를 통한 예측을했다면 예측 시간을 훨씬 단축시키진 않았을까 라는 생각이 들었다. 실험을통해 softmax를 사용했을 때 SVM보다 성능이 떨어졌다고는 했지만, 하나의 분류기를 사용하는 것을 중점적으로 실험을 했다면 조금 더 간단한 분류기가 구축되었을 것이라는 의견이다. 물론, 이미 많은 실험들을 통해 SVM을 선택한 것이지만.</li>
</ul>

<p><br /><br /></p>

<h1 id="구현">구현</h1>
<p>Pytorch로 RCNN을 구현해보자. 데이터는 Roboflow에서 제공하는 <a href="https://universe.roboflow.com/yinguo/soccer-data">soccer dataset</a>을 사용했다. Class는 볼을 소유한 player와 소유하지않은
player 주로 이 2개로 이루어져있지만, coco json 상에서는 배경과 일반 player class가 존재해 총 4개의 클래스(0: players, 1: None, 2: has ball, 3: no ball)로 명시되어있다.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="err">!</span><span class="n">pip</span> <span class="n">install</span> <span class="o">-</span><span class="n">q</span> <span class="n">selectivesearch</span> <span class="c1"># selective search 라이브러리
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">json</span>
<span class="kn">import</span> <span class="nn">cv2</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">os</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 데이터 경로 설정
</span><span class="n">DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/'</span>
<span class="n">TR_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/train/'</span>
<span class="n">VAL_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/valid/'</span>
<span class="n">TEST_DATA_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/test/'</span>
<span class="n">TR_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/train/_annotations.coco.json'</span>
<span class="n">VAL_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/valid/_annotations.coco.json'</span>
<span class="n">TEST_LAB_PATH</span> <span class="o">=</span> <span class="s">'/content/drive/MyDrive/논문실습/data/coco_format/test/_annotations.coco.json'</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># annotaions file
</span><span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">TR_LAB_PATH</span><span class="p">,</span> <span class="s">'r'</span><span class="p">)</span> <span class="k">as</span> <span class="n">f</span><span class="p">:</span>
    <span class="n">tr_lab</span> <span class="o">=</span> <span class="n">json</span><span class="p">.</span><span class="n">load</span><span class="p">(</span><span class="n">f</span><span class="p">)</span>
<span class="c1">#print(json.dumps(tr_lab, indent=4))
</span>
<span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">TEST_LAB_PATH</span><span class="p">,</span> <span class="s">'r'</span><span class="p">)</span> <span class="k">as</span> <span class="n">f</span><span class="p">:</span>
    <span class="n">test_lab</span> <span class="o">=</span> <span class="n">json</span><span class="p">.</span><span class="n">load</span><span class="p">(</span><span class="n">f</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 이미지 파일목록
</span><span class="n">tr_images</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">TR_DATA_PATH</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.jpg'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>
<span class="n">val_images</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">VAL_DATA_PATH</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.jpg'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>
<span class="nb">len</span><span class="p">(</span><span class="n">tr_images</span><span class="p">),</span> <span class="nb">len</span><span class="p">(</span><span class="n">val_images</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># IoU 계산 함수
</span><span class="k">def</span> <span class="nf">get_iou</span><span class="p">(</span><span class="n">cand_box</span><span class="p">,</span> <span class="n">gt_box</span><span class="p">):</span> <span class="c1"># cand_box: selective search box, gt_box: ground_truth
</span>    <span class="c1"># cand box는 left, top, right, bottom 순
</span>    <span class="k">assert</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span>
    <span class="k">assert</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span>
    <span class="k">assert</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'x1'</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'x2'</span><span class="p">]</span>
    <span class="k">assert</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'y1'</span><span class="p">]</span> <span class="o">&lt;</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'y2'</span><span class="p">]</span>
    <span class="n">x_left</span> <span class="o">=</span> <span class="nb">max</span><span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'x1'</span><span class="p">])</span> <span class="c1"># 겹치는 구간을 파악하기 위해 left는 더 큰 것
</span>    <span class="n">y_top</span> <span class="o">=</span> <span class="nb">max</span><span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'y1'</span><span class="p">])</span>
    <span class="n">x_right</span> <span class="o">=</span> <span class="nb">min</span><span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'x2'</span><span class="p">])</span> <span class="c1"># right는 둘 중 더 작은 것
</span>    <span class="n">y_bottom</span> <span class="o">=</span> <span class="nb">min</span><span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">3</span><span class="p">],</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'y2'</span><span class="p">])</span>
    <span class="k">if</span> <span class="n">x_right</span> <span class="o">&lt;</span> <span class="n">x_left</span> <span class="ow">or</span> <span class="n">y_bottom</span> <span class="o">&lt;</span> <span class="n">y_top</span><span class="p">:</span> <span class="c1"># 겹치지 않는 경우
</span>        <span class="k">return</span> <span class="mf">0.0</span>
    <span class="n">intersection_area</span> <span class="o">=</span> <span class="p">(</span><span class="n">x_right</span> <span class="o">-</span> <span class="n">x_left</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">y_bottom</span> <span class="o">-</span> <span class="n">y_top</span><span class="p">)</span>
    <span class="n">bb1_area</span> <span class="o">=</span> <span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span> <span class="o">-</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span> <span class="o">*</span> <span class="p">(</span><span class="n">cand_box</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span> <span class="o">-</span> <span class="n">cand_box</span><span class="p">[</span><span class="mi">1</span><span class="p">])</span>
    <span class="n">bb2_area</span> <span class="o">=</span> <span class="p">(</span><span class="n">gt_box</span><span class="p">[</span><span class="s">'x2'</span><span class="p">]</span> <span class="o">-</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'x1'</span><span class="p">])</span> <span class="o">*</span> <span class="p">(</span><span class="n">gt_box</span><span class="p">[</span><span class="s">'y2'</span><span class="p">]</span> <span class="o">-</span> <span class="n">gt_box</span><span class="p">[</span><span class="s">'y1'</span><span class="p">])</span>
    <span class="n">iou</span> <span class="o">=</span> <span class="n">intersection_area</span> <span class="o">/</span> <span class="nb">float</span><span class="p">(</span><span class="n">bb1_area</span> <span class="o">+</span> <span class="n">bb2_area</span> <span class="o">-</span> <span class="n">intersection_area</span><span class="p">)</span>
    <span class="k">assert</span> <span class="n">iou</span> <span class="o">&gt;=</span> <span class="mf">0.0</span>
    <span class="k">assert</span> <span class="n">iou</span> <span class="o">&lt;=</span> <span class="mf">1.0</span>
    <span class="k">return</span> <span class="n">iou</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 파일명을 넣으면 img id가 리턴되도록 dictionary 만들기
</span><span class="n">img_id_dict</span> <span class="o">=</span> <span class="nb">dict</span><span class="p">()</span>
<span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">tr_lab</span><span class="p">[</span><span class="s">'images'</span><span class="p">]:</span>
    <span class="n">img_id_dict</span><span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'file_name'</span><span class="p">]]</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="s">'id'</span><span class="p">]</span>

<span class="c1"># 파일명을 넣으면 test img id가 리턴되도록 dictionary 만들기
</span><span class="n">test_img_id_dict</span> <span class="o">=</span> <span class="nb">dict</span><span class="p">()</span>
<span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">test_lab</span><span class="p">[</span><span class="s">'images'</span><span class="p">]:</span>
    <span class="n">test_img_id_dict</span><span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'file_name'</span><span class="p">]]</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="s">'id'</span><span class="p">]</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># selective search 결과를 전달받아 iou를 계산하고 positive, negative로 분류하여
# 결과를 리턴해주는 함수. RCNN은 이미지당 2000개의 region을 제안하지만 GPU의 한계로
# positive, negative 각각 30개씩만 제안받도록함.
</span><span class="k">def</span> <span class="nf">pos_neg_region</span><span class="p">(</span><span class="n">image</span><span class="p">,</span> <span class="n">ssresults</span><span class="p">,</span> <span class="n">bboxes</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">threshold</span><span class="p">):</span>
    <span class="n">train_imgs</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">train_labs</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">train_cands</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="n">p_num</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="n">n_num</span> <span class="o">=</span> <span class="mi">0</span>
    <span class="n">cands</span> <span class="o">=</span> <span class="p">[</span><span class="n">cand</span><span class="p">[</span><span class="s">'rect'</span><span class="p">]</span> <span class="k">for</span> <span class="n">cand</span> <span class="ow">in</span> <span class="n">ssresults</span> <span class="k">if</span> <span class="n">cand</span><span class="p">[</span><span class="s">'size'</span><span class="p">]</span> <span class="o">&lt;</span> <span class="mi">10000</span><span class="p">]</span> <span class="c1"># selective search box 중에 크기가 10000이하인 box 좌표만
</span>    <span class="n">cand_rects</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">cands</span><span class="p">:</span>
        <span class="k">if</span> <span class="n">x</span> <span class="ow">not</span> <span class="ow">in</span> <span class="n">cand_rects</span><span class="p">:</span>
            <span class="n">cand_rects</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="c1"># 영역 중복 제거
</span>
<span class="c1">#    print(len(cand_rects))
</span>
    <span class="k">for</span> <span class="n">_</span><span class="p">,</span> <span class="n">cand</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">cand_rects</span><span class="p">):</span>
        <span class="n">cand</span> <span class="o">=</span> <span class="nb">list</span><span class="p">(</span><span class="n">cand</span><span class="p">)</span>
        <span class="k">if</span> <span class="n">cand</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># 넓이, 높이가 0이면 박스가 그려지지않기때문에 1로 바꿈
</span>            <span class="n">cand</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>

        <span class="k">elif</span> <span class="n">cand</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
            <span class="n">cand</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>
        <span class="n">cand</span><span class="p">[</span><span class="mi">2</span><span class="p">]</span> <span class="o">+=</span> <span class="n">cand</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="c1"># width -&gt; x2 형식.
</span>        <span class="n">cand</span><span class="p">[</span><span class="mi">3</span><span class="p">]</span> <span class="o">+=</span> <span class="n">cand</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="c1"># height -&gt; y2 형식.
</span>
        <span class="k">if</span> <span class="n">p_num</span> <span class="o">&gt;</span> <span class="mi">30</span> <span class="ow">and</span> <span class="n">n_num</span> <span class="o">&gt;</span> <span class="mi">30</span><span class="p">:</span> <span class="c1"># gpu의 한계로 positive, negative 각각 30개씩만. region이 2000개면 너무 많음 
</span>            <span class="k">break</span>

        <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="nb">len</span><span class="p">(</span><span class="n">labels</span><span class="p">)):</span>
            <span class="n">bbox</span> <span class="o">=</span> <span class="n">bboxes</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>
            <span class="n">label</span> <span class="o">=</span> <span class="n">labels</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>
            <span class="n">img</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">resize</span><span class="p">(</span><span class="n">image</span><span class="p">[</span><span class="n">cand</span><span class="p">[</span><span class="mi">1</span><span class="p">]:</span><span class="n">cand</span><span class="p">[</span><span class="mi">3</span><span class="p">],</span> <span class="n">cand</span><span class="p">[</span><span class="mi">0</span><span class="p">]:</span><span class="n">cand</span><span class="p">[</span><span class="mi">2</span><span class="p">]],</span> <span class="p">(</span><span class="mi">224</span><span class="p">,</span> <span class="mi">224</span><span class="p">),</span> <span class="n">interpolation</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">INTER_CUBIC</span><span class="p">)</span> <span class="c1"># 원본에서 후보 부분 잘라서 저장
</span>            <span class="n">iou</span> <span class="o">=</span> <span class="n">get_iou</span><span class="p">(</span><span class="n">cand</span><span class="p">,</span> <span class="p">{</span><span class="s">"x1"</span><span class="p">:</span><span class="n">bbox</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span><span class="s">"x2"</span><span class="p">:</span><span class="n">bbox</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span><span class="o">+</span><span class="n">bbox</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span><span class="s">"y1"</span><span class="p">:</span><span class="n">bbox</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span><span class="s">"y2"</span><span class="p">:</span><span class="n">bbox</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">+</span> <span class="n">bbox</span><span class="p">[</span><span class="mi">3</span><span class="p">]})</span>

            <span class="k">if</span> <span class="n">iou</span> <span class="o">&gt;</span> <span class="n">threshold</span><span class="p">:</span>
                <span class="k">if</span> <span class="n">p_num</span> <span class="o">&lt;</span> <span class="mi">30</span><span class="p">:</span>
                    <span class="c1">#print(cand, bbox)
</span>                    <span class="n">train_imgs</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">img</span><span class="p">)</span>
                    <span class="n">train_labs</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="nb">int</span><span class="p">(</span><span class="n">label</span><span class="p">))</span>
                    <span class="n">train_cands</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">cand</span><span class="p">)</span>
                    <span class="n">p_num</span> <span class="o">+=</span> <span class="mi">1</span>

            <span class="k">else</span><span class="p">:</span>
                <span class="k">if</span> <span class="n">n_num</span> <span class="o">&lt;</span> <span class="mi">30</span><span class="p">:</span>
                    <span class="n">train_imgs</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">img</span><span class="p">)</span>
                    <span class="n">train_labs</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>
                    <span class="n">train_cands</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">cand</span><span class="p">)</span>
                    <span class="n">n_num</span> <span class="o">+=</span> <span class="mi">1</span>

    <span class="k">return</span> <span class="n">train_imgs</span><span class="p">,</span> <span class="n">train_labs</span><span class="p">,</span> <span class="n">train_cands</span>

</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">selectivesearch</span>

<span class="k">def</span> <span class="nf">region_proposal</span><span class="p">(</span><span class="n">image</span><span class="p">,</span> <span class="n">mode</span><span class="p">):</span>
    <span class="n">train_images</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">train_labels</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">train_cands</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="k">if</span> <span class="n">mode</span> <span class="o">==</span> <span class="s">'finetuning'</span><span class="p">:</span>
        <span class="n">threshold</span> <span class="o">=</span> <span class="mf">0.3</span> <span class="c1"># 논문처럼 0.5로 했더니 positive region이 잘 탐색되지 않음
</span>        <span class="n">img</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">imread</span><span class="p">(</span><span class="n">TR_DATA_PATH</span> <span class="o">+</span> <span class="n">image</span><span class="p">)</span>
        <span class="nb">id</span> <span class="o">=</span> <span class="n">img_id_dict</span><span class="p">[</span><span class="n">image</span><span class="p">]</span>
        <span class="n">bboxes</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'bbox'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">tr_lab</span><span class="p">[</span><span class="s">'annotations'</span><span class="p">]</span> <span class="k">if</span> <span class="n">x</span><span class="p">[</span><span class="s">'image_id'</span><span class="p">]</span> <span class="o">==</span> <span class="nb">id</span><span class="p">]</span>
        <span class="n">labels</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'category_id'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">tr_lab</span><span class="p">[</span><span class="s">'annotations'</span><span class="p">]</span> <span class="k">if</span> <span class="n">x</span><span class="p">[</span><span class="s">'image_id'</span><span class="p">]</span> <span class="o">==</span> <span class="nb">id</span><span class="p">]</span>
    <span class="k">elif</span> <span class="n">mode</span> <span class="o">==</span> <span class="s">'classify'</span><span class="p">:</span>
        <span class="n">threshold</span> <span class="o">=</span> <span class="mf">0.3</span>
    <span class="k">elif</span> <span class="n">mode</span> <span class="o">==</span> <span class="s">'test'</span><span class="p">:</span>
        <span class="n">threshold</span> <span class="o">=</span> <span class="mf">0.3</span>
        <span class="n">img</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">imread</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span> <span class="o">+</span> <span class="n">image</span><span class="p">)</span>
        <span class="nb">id</span> <span class="o">=</span> <span class="n">test_img_id_dict</span><span class="p">[</span><span class="n">image</span><span class="p">]</span>
        <span class="n">bboxes</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'bbox'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">test_lab</span><span class="p">[</span><span class="s">'annotations'</span><span class="p">]</span> <span class="k">if</span> <span class="n">x</span><span class="p">[</span><span class="s">'image_id'</span><span class="p">]</span> <span class="o">==</span> <span class="nb">id</span><span class="p">]</span>
        <span class="n">labels</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span><span class="p">[</span><span class="s">'category_id'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">test_lab</span><span class="p">[</span><span class="s">'annotations'</span><span class="p">]</span> <span class="k">if</span> <span class="n">x</span><span class="p">[</span><span class="s">'image_id'</span><span class="p">]</span> <span class="o">==</span> <span class="nb">id</span><span class="p">]</span>

    <span class="c1"># 이미지당 gt box 만들기
</span>    <span class="n">_</span><span class="p">,</span> <span class="n">regions</span> <span class="o">=</span> <span class="n">selectivesearch</span><span class="p">.</span><span class="n">selective_search</span><span class="p">(</span><span class="n">img</span><span class="p">,</span> <span class="n">scale</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">min_size</span><span class="o">=</span><span class="mi">100</span><span class="p">)</span>
    <span class="n">imgs</span><span class="p">,</span> <span class="n">labs</span><span class="p">,</span> <span class="n">cands</span> <span class="o">=</span> <span class="n">pos_neg_region</span><span class="p">(</span><span class="n">img</span><span class="p">,</span> <span class="n">regions</span><span class="p">,</span> <span class="n">bboxes</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span>  <span class="n">threshold</span><span class="p">)</span>

    <span class="n">train_images</span> <span class="o">+=</span> <span class="n">imgs</span> <span class="c1"># 한 이미지당 selectivesearch 결과가 저장된 리스트
</span>    <span class="n">train_labels</span> <span class="o">+=</span> <span class="n">labs</span>
    <span class="n">train_cands</span> <span class="o">+=</span> <span class="n">cands</span>

    <span class="k">if</span> <span class="n">mode</span> <span class="o">==</span> <span class="s">'test'</span><span class="p">:</span>
      <span class="k">return</span> <span class="n">train_images</span><span class="p">,</span> <span class="n">train_labels</span><span class="p">,</span> <span class="n">train_cands</span> <span class="c1"># test시에는 시각화를 위해 후보 영역의 좌표 정보도 같이 리턴
</span>    <span class="k">return</span> <span class="n">train_images</span><span class="p">,</span> <span class="n">train_labels</span>

</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 모델 정의
</span><span class="kn">from</span> <span class="nn">torchvision.models</span> <span class="kn">import</span> <span class="n">mobilenet_v3_small</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch</span>

<span class="c1"># Mobilenet을 CNN으로 사용
</span><span class="n">model</span> <span class="o">=</span> <span class="n">mobilenet_v3_small</span><span class="p">(</span><span class="n">num_classes</span><span class="o">=</span><span class="mi">4</span><span class="p">)</span> <span class="c1"># 배경, player, 공을 소유한 player, 공을 소유하지 않은 player
# coco 형식 상에서는 4개의 클래스지만 사실상 공을 소유한 player와 공을 소유하지 않은 player인 2개의 클래스로 구분된다.
</span>
<span class="n">params</span> <span class="o">=</span> <span class="p">[</span><span class="n">p</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">parameters</span><span class="p">()</span> <span class="k">if</span> <span class="n">p</span><span class="p">.</span><span class="n">requires_grad</span><span class="p">]</span>

<span class="c1"># linear SVM을 직접구현하는대신 편의상 pytorch에서 제공하는 선형 분류기 사용
</span><span class="n">model</span><span class="p">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">576</span><span class="p">,</span> <span class="mi">4096</span><span class="p">),</span> <span class="c1"># mobilenet 마지막 layer의 output576
</span>                    <span class="n">nn</span><span class="p">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">4096</span><span class="p">,</span> <span class="mi">4</span><span class="p">))</span>  <span class="c1"># 배경포함
</span><span class="n">device</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cuda'</span><span class="p">)</span> <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">cuda</span><span class="p">.</span><span class="n">is_available</span><span class="p">()</span> <span class="k">else</span> <span class="n">torch</span><span class="p">.</span><span class="n">device</span><span class="p">(</span><span class="s">'cpu'</span><span class="p">)</span>
<span class="n">model</span><span class="p">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 손실함수 및 optimizer 정의
</span><span class="kn">import</span> <span class="nn">torch.optim</span> <span class="k">as</span> <span class="n">optim</span>

<span class="n">criterion</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">CrossEntropyLoss</span><span class="p">().</span><span class="n">cuda</span><span class="p">()</span>
<span class="n">optimizer</span> <span class="o">=</span> <span class="n">optim</span><span class="p">.</span><span class="n">SGD</span><span class="p">(</span><span class="n">params</span><span class="p">,</span><span class="n">lr</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span> <span class="n">momentum</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">weight_decay</span><span class="o">=</span><span class="mf">0.0005</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 정규화 및 데이터 증강
# 정규화를 하지않으면 gradient exploding 일어날 수 있음
</span><span class="kn">from</span> <span class="nn">torchvision</span> <span class="kn">import</span> <span class="n">transforms</span>
<span class="n">train_transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="p">.</span><span class="n">Compose</span><span class="p">([</span>
                                      <span class="n">transforms</span><span class="p">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="c1"># 0과 1 사이값을 가지도록 정규화
</span>                                      <span class="n">transforms</span><span class="p">.</span><span class="n">RandomVerticalFlip</span><span class="p">(</span><span class="n">p</span><span class="o">=</span><span class="mf">0.5</span><span class="p">),</span>
                                      <span class="n">transforms</span><span class="p">.</span><span class="n">RandomHorizontalFlip</span><span class="p">(</span><span class="n">p</span><span class="o">=</span><span class="mf">0.5</span><span class="p">),</span>
<span class="p">])</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 학습
</span><span class="kn">import</span> <span class="nn">time</span>

<span class="n">num_epochs</span> <span class="o">=</span> <span class="mi">50</span>
<span class="k">print</span><span class="p">(</span><span class="s">'----------------------train start--------------------------'</span><span class="p">)</span>

<span class="k">for</span> <span class="n">epoch</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">num_epochs</span><span class="p">):</span>

  <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span>
  <span class="n">model</span><span class="p">.</span><span class="n">train</span><span class="p">()</span>
  <span class="n">epoch_loss</span> <span class="o">=</span> <span class="mi">0</span>
  <span class="n">prog_bar</span> <span class="o">=</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">train_data</span><span class="p">,</span> <span class="n">total</span><span class="o">=</span><span class="nb">len</span><span class="p">(</span><span class="n">train_data</span><span class="p">))</span>

  <span class="k">for</span> <span class="n">_</span><span class="p">,</span> <span class="n">t</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">prog_bar</span><span class="p">):</span>
      <span class="n">image_data</span><span class="p">,</span> <span class="n">label_data</span> <span class="o">=</span> <span class="n">t</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">t</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="c1"># 원래는 box regression으로 region을 조정 후 조정된 region 입력값이 들어와야함
</span>      <span class="c1"># 논문에서는 positive 32, negative 96개를 포함하는 128의 batch size사용. 편의상 하나의 이미지 내 region의 갯수를 batch size로 사용.
</span>      <span class="n">inputs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">(</span><span class="nb">tuple</span><span class="p">(</span><span class="n">train_transform</span><span class="p">(</span><span class="nb">id</span><span class="p">).</span><span class="n">cuda</span><span class="p">().</span><span class="n">reshape</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">224</span><span class="p">,</span> <span class="mi">224</span><span class="p">)</span> <span class="k">for</span> <span class="nb">id</span> <span class="ow">in</span> <span class="n">image_data</span><span class="p">))</span> <span class="c1"># input을 (?,3,224,224) 사이즈로 바꿈
</span>      <span class="n">labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">(</span><span class="n">label_data</span><span class="p">).</span><span class="n">cuda</span><span class="p">()</span>
      <span class="n">labels</span> <span class="o">=</span> <span class="n">labels</span><span class="p">.</span><span class="nb">type</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="nb">long</span><span class="p">)</span>

      <span class="n">outputs</span> <span class="o">=</span> <span class="n">model</span><span class="p">(</span><span class="n">inputs</span><span class="p">)</span> <span class="c1"># (region, 각 클래스에 속하는 출력값)의 크기를 갖는 텐서
</span>
      <span class="c1"># 만약 outputs이 0과 1사이인 확률값이 아니라면 crossentropy를 계산할 때 자동으로 확률값으로 변환 후 손실 계산
</span>      <span class="c1"># 각 region마다 4개의 클래스에 대한 출력값을 확률값으로 변환시켜 손실 계산
</span>      <span class="n">loss</span> <span class="o">=</span> <span class="n">criterion</span><span class="p">(</span><span class="n">outputs</span><span class="p">,</span> <span class="n">labels</span><span class="p">)</span> <span class="c1"># loss에러나면 모델 빌딩부터 다시
#      print(loss)
</span>      <span class="n">loss</span><span class="p">.</span><span class="n">backward</span><span class="p">()</span>
      <span class="n">optimizer</span><span class="p">.</span><span class="n">step</span><span class="p">()</span>
      <span class="n">epoch_loss</span> <span class="o">+=</span> <span class="n">loss</span><span class="p">.</span><span class="n">item</span><span class="p">()</span>
  <span class="k">print</span><span class="p">(</span><span class="sa">f</span><span class="s">'epoch : </span><span class="si">{</span><span class="n">epoch</span><span class="o">+</span><span class="mi">1</span><span class="si">}</span><span class="s">, Loss : </span><span class="si">{</span><span class="n">epoch_loss</span><span class="si">}</span><span class="s">, time : </span><span class="si">{</span><span class="n">time</span><span class="p">.</span><span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># model test
</span><span class="n">test_images</span> <span class="o">=</span> <span class="p">[</span><span class="n">x</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">os</span><span class="p">.</span><span class="n">listdir</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span><span class="p">))</span> <span class="k">if</span> <span class="s">'.jpg'</span> <span class="ow">in</span> <span class="n">x</span><span class="p">]</span>

<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>

<span class="n">candidate_predict</span> <span class="o">=</span> <span class="p">[]</span>
<span class="n">candidate_score</span> <span class="o">=</span> <span class="p">[]</span>

<span class="n">model</span><span class="p">.</span><span class="nb">eval</span><span class="p">()</span>
<span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="n">no_grad</span><span class="p">():</span>
      <span class="c1"># 시각화를 위해 region의 좌표정보도 함께 가져옴
</span>      <span class="c1"># test용으로 이미지 하나만 실행
</span>      <span class="n">image_data</span><span class="p">,</span> <span class="n">label_data</span><span class="p">,</span> <span class="n">regions</span> <span class="o">=</span> <span class="n">region_proposal</span><span class="p">(</span><span class="n">test_images</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">mode</span> <span class="o">=</span> <span class="s">'test'</span><span class="p">)</span>
      <span class="n">inputs</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">(</span><span class="nb">tuple</span><span class="p">(</span><span class="n">train_transform</span><span class="p">(</span><span class="nb">id</span><span class="p">).</span><span class="n">cuda</span><span class="p">().</span><span class="n">reshape</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">224</span><span class="p">,</span> <span class="mi">224</span><span class="p">)</span> <span class="k">for</span> <span class="nb">id</span> <span class="ow">in</span> <span class="n">image_data</span><span class="p">))</span> <span class="c1"># input을 (?,3,224,224) 사이즈로 바꿈
</span>      <span class="n">labels</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">Tensor</span><span class="p">(</span><span class="n">label_data</span><span class="p">).</span><span class="n">cuda</span><span class="p">()</span>
      <span class="n">labels</span> <span class="o">=</span> <span class="n">labels</span><span class="p">.</span><span class="nb">type</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="nb">long</span><span class="p">)</span>

      <span class="n">outputs</span> <span class="o">=</span> <span class="n">model2</span><span class="p">(</span><span class="n">inputs</span><span class="p">)</span>
      <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
        <span class="n">predict_class</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">).</span><span class="n">item</span><span class="p">()</span> <span class="c1"># 예측한 클래스. 출력값이 가장 높은 값
</span>        <span class="n">candidate_predict</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">predict_class</span><span class="p">)</span>
        <span class="n">candidate_score</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">F</span><span class="p">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">x</span><span class="p">)[</span><span class="n">predict_class</span><span class="p">].</span><span class="n">item</span><span class="p">())</span> <span class="c1"># 해당 클래스의 확률값
</span></code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 여러 region이 겹칠 수 있으니까 가장 정답과 가까우면서 다른 region과 겹치지않는 박스만 남겨두기
# test시에는 nms가 적용되어 정리된 결과만을 출력해야한다.
</span>
<span class="kn">from</span> <span class="nn">torchvision.ops</span> <span class="kn">import</span> <span class="n">nms</span>

<span class="c1"># IoU가 0.2이상이면 겹치는 것으로 간주하고 제거해서 가장 score가 높은 박스만 남겨두기
</span><span class="n">selected_idx</span> <span class="o">=</span> <span class="n">nms</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">regions</span><span class="p">).</span><span class="nb">float</span><span class="p">(),</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">candidate_score</span><span class="p">),</span> <span class="n">iou_threshold</span> <span class="o">=</span> <span class="mf">0.2</span><span class="p">)</span>
<span class="n">selected_boxes</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">regions</span><span class="p">)[</span><span class="n">selected_idx</span><span class="p">]</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># 원본 이미지 시각화
</span><span class="n">img_test</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">imread</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span> <span class="o">+</span> <span class="n">test_images</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span>
<span class="n">img_test</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">cvtColor</span><span class="p">(</span><span class="n">img_test</span><span class="p">,</span> <span class="n">cv2</span><span class="p">.</span><span class="n">COLOR_BGR2RGB</span><span class="p">)</span>

<span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">outputs</span><span class="p">):</span>
  <span class="k">if</span> <span class="n">torch</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">!=</span> <span class="mi">0</span><span class="p">:</span>
    <span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="p">,</span><span class="n">x2</span><span class="p">,</span><span class="n">y2</span> <span class="o">=</span> <span class="n">regions</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">img_test</span><span class="p">,</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="p">),</span> <span class="p">(</span><span class="n">x2</span><span class="p">,</span><span class="n">y2</span><span class="p">),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">img_test</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">img_test</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="374" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/89bcb5d7-07b6-4a15-ab29-66e41d8056b2" />
  </p>
</div>

<p><br /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># nms로 선택된 이미지 시각화
</span><span class="n">img_test_nms</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">imread</span><span class="p">(</span><span class="n">TEST_DATA_PATH</span> <span class="o">+</span> <span class="n">test_images</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span>
<span class="n">img_test_nms</span> <span class="o">=</span> <span class="n">cv2</span><span class="p">.</span><span class="n">cvtColor</span><span class="p">(</span><span class="n">img_test_nms</span><span class="p">,</span> <span class="n">cv2</span><span class="p">.</span><span class="n">COLOR_BGR2RGB</span><span class="p">)</span>

<span class="k">for</span> <span class="n">idx</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">selected_boxes</span><span class="p">):</span>
    <span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="p">,</span><span class="n">x2</span><span class="p">,</span><span class="n">y2</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">tolist</span><span class="p">()</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">rectangle</span><span class="p">(</span><span class="n">img_test_nms</span><span class="p">,</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="p">),</span> <span class="p">(</span><span class="n">x2</span><span class="p">,</span><span class="n">y2</span><span class="p">),</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>
    <span class="n">cv2</span><span class="p">.</span><span class="n">putText</span><span class="p">(</span><span class="n">img_test_nms</span><span class="p">,</span> <span class="nb">str</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">x</span><span class="p">).</span><span class="n">tolist</span><span class="p">()),</span> <span class="p">(</span><span class="n">x1</span><span class="p">,</span><span class="n">y1</span><span class="o">-</span><span class="mi">10</span><span class="p">),</span> <span class="n">cv2</span><span class="p">.</span><span class="n">FONT_HERSHEY_SIMPLEX</span><span class="p">,</span> <span class="mf">0.7</span><span class="p">,</span> <span class="n">color</span> <span class="o">=</span> <span class="p">(</span><span class="mi">255</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">thickness</span><span class="o">=</span> <span class="mi">3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">imshow</span><span class="p">(</span><span class="n">img_test_nms</span><span class="p">)</span>
</code></pre></div></div>

<p><br /></p>
<div align="center">
  <p>
  <img width="374" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/14ffddf3-6d7e-4c20-b7c6-7b0f6e6ca5ed" />
  </p>
</div>

<p><br /></p>

<p><br /></p>

<p>위의 이미지를 보면 localization이 잘 수행되지않아 box가 객체를 잘 감싸지 못하고 있다. RCNN은 위의 코드와 같이 region proposals과 classification을 수행할 수 있고, 더 정확한 localization을 위해 box regression을 독립적으로 수행해야하지만 일단 다음을 기약한다… Box regression 후에는 더 정확한 예측을 수행할 것이다.</p>

<p>실제로해보니 region proposals, classification, box regression이 모두 독립적으로 수행되니까 굉장히 비효율적이고 한 stage마다 너무 많은 시간이 소요된다. 이후에 등장한 one-stage model이나
Faster RCNN과 같은 모델에 감사하자.</p>

<h1 id="이미지-출처">이미지 출처</h1>
<ul>
  <li><a href="https://arxiv.org/pdf/1311.2524.pdf">Rich feature hierarchies for accurate object detection and semantic segmentation Paper</a></li>
  <li><a href="https://wikidocs.net/142645">NMS Image</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Detection" /><summary type="html"><![CDATA[Abstract 지난 몇년동안(2013) PASCAL VOC dataset에 관해 객체 탐지의 성능 향상이 더디다. 대화마다 best model은 일반적으로 여러개의 저수준 피처와 고수준의 context를 결합한 복잡한 앙상블 구조이다. 본 연구에서는 단순한 탐지 알고리즘으로 VOC 2012에서 mAP 53.3%를 달성하며 지난 대회 대비 30% 이상 성능을 향상시킨 방법에 대해 제안한다. 본 논문의 두가지 키포인트는 첫째, bottom-up region proposal에 high capacity CNN을 적용해서 localize와 segment objects를 할 수 있는 것과 둘째, 훈련 데이터가 부족한 경우 사전 학습 모델을 fine tuning하여 성능을 크게 높였다는 것이다. Region proposal이 CNN과 결합되었기 때문에 R-CNN이라 부르고, 최근 개발된 sliding window 개념을 활용한 CNN 아키텍처인 OverFeat과 비교했을 때 ILSVRC2013 탐지부문에서 큰 격차로 우승했다.]]></summary></entry><entry><title type="html">DeepLabv3+ : Encoder-Decoder with Atrous Separable Convolution for Semantic Image Segmentation 논문 요약</title><link href="https://kimyeonz.github.io/blog/segmentation/DeepLabV3_Plus/" rel="alternate" type="text/html" title="DeepLabv3+ : Encoder-Decoder with Atrous Separable Convolution for Semantic Image Segmentation 논문 요약" /><published>2023-10-10T00:00:00+00:00</published><updated>2023-10-10T00:00:00+00:00</updated><id>https://kimyeonz.github.io/blog/segmentation/DeepLabV3_Plus</id><content type="html" xml:base="https://kimyeonz.github.io/blog/segmentation/DeepLabV3_Plus/"><![CDATA[<h1 id="atrous-convolution이란">Atrous Convolution이란?</h1>
<p>DeepLabv3+ 논문을 리뷰하기 앞서 DeepLab 모델들에서 사용된 핵심 기법인 Atrous Convolution에 대해 알아보자.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/f677c6c9-c341-4081-bba7-cb7c8d6218f1" />
  </p>
  <p>3x3 Atrous covolution, stride 1, padding = rate - 1 일 때 출력되는 피처맵</p>
</div>

<p><br /><br /></p>

<p>Atrous convolution은 구멍이라는 말처럼 convolution 단계에서 필터 사이사이에 0을 입력하여 간격을 두는 기법으로 더 넓은 receptive field를 가질 수 있게 한다. 여기서 recpetive field란, 필터의 한 노드가 입력으로부터 수용할 수 있는 영역으로 segmentation과 같이 분류에 비해 정밀한 task가 요구된다면 주변의 공간 정보를 넓게 잘 활용해야하므로 넓은 receptive field를 가지는 것이 좋다.</p>

<p>넓은 receptive field를 위해 더 사이즈가 큰 필터를 사용하면 되지 않느냐는 생각이 들 수 있는데 필터의 사이즈가 커지면 파라미터의 갯수도 증가하기때문에 파라미터의 갯수는 그대로 유지시키면서 넓은 receptive field의 효과를 낼 수 있는 방법이 atrous(dilated) convolution이다.</p>

<p>위의 그림에서 빨간색 점을 필터의 크기, dilation rate을 $2^n$씩 증가시킨다고 할 때 (a)는 dilation rate이 1이기때문에 일반적인 convolution과 같다. (b)는 dilation rate이 2이고, padding은 1이기때문에 수용 영역이 7x7이 된다. 같은 방법으로 (c)는 dilation rate이 4이고, padding은 3이기때문에 15x15의 receptive field를 가진다. 결과적으로 3x3이라는 같은 필터의 크기를 가졌음에도 rate에 따라 필터사이에 0을 집어넣어 연산량은 일반적인 3x3 합성곱과 똑같이 가져가면서 padding까지 겹쳐져 훨씬 더 큰 receptive field를 가지는 것을 알 수 있다.</p>

<p>Receptive field를 구하는 공식은 다음과 같다.</p>

<p><br /></p>

\[ReceptiveField = ((Filter Size - 1) * Dilated Rate) + 1 +\]

\[　　　　　　　　　　　　(Filter Size  - 1) * (Stride - 1) + 2 * Padding\]

<p><br /> 
추가적으로, receptive field는 layer가 깊어질수록 커진다. Receptive field란, 필터의 한 노드가 입력으로부터 수용할 수 있는 영역의 크기라고 했는데 아래 그림과 같이 층이 깊어질수록 한 노드가 갖는 수용영역의 크기는 넓어질 것이다. layer2 필터의 한 노드는 3x3의 receptive field를 가지고, layer3 필터의 한 노드는 layer 2의 3x3 -&gt; 결국 layer 1의 5x5 모든 영역을 receptive field로 가지는 것이 된다. 층이 깊어질수록 receptive field가 커지고 이로 인해 저수준의 특징에서 점점 고수준의 특징을 학습할 수 있게된다.</p>

<p><br /></p>
<div align="center">
  <p>
  <img width="500" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/59cfdef6-42f5-4d31-a34e-6e42700dce6e" />
  </p>
</div>

<p><br /></p>

<p>아래 그림을 보면 똑같은 필터 크기를 사용하더라도 일반적인 convolution을 사용해 추출한 sparse한 features보다 atrous convolution 기법을 적용한 dense한 features의 object가 더 선명하게 보이는 것을 알 수 있다.</p>

<div align="center">
  <p>
  <img width="600" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/57d8afaa-2932-4f92-99a4-f80f01334fb3" />
  </p>
</div>

<p><br /><br /></p>

<p>Classification task에서는 한 이미지 내에서 object가 존재하는지의 여부가 중요했다면 semantic segmentation task에서는 단순히 object의 존재여부보다 object 간의 경계를 잘 파악하는 훨씬 중요하다. 
따라서, classification model은 receptive field가 상대적으로 작고 분류기에 fully connected layer를 통해 공간정보를 많이 손실하기 때문에 object의 경계간 detail한 정보를 얻기가 힘들어서 deeplab에서는 이를 해결하기 위해 atrous convolution을 사용하여 receptive field를 키우고 분류기에는 U-net처럼 fully convolutional layer를 사용해 공간정보를 최대한 활용하여 segmentation task를 잘 수행했다.</p>

<h1 id="abstract">Abstract</h1>
<p>Spatial pyramid pooling module이나 encode-decoder 구조가 semantic segmentation 분야에서 잘 사용되고 있다. Spatial pyramid pooling 구조는 필터의 크기를 크게해서 넓은 영역의 context information을 encoding할 수 있고, encode-decoder 구조는 network의 공간정보를 점진적으로 복구해서 object 간의 경계를 더 선명하게 포착할 수 있게 한다. 본 논문에서는 이 두가지 방법을 모두 활용한 DeepLabv3+ 모델에 대해 설명한다. 추가적으로 Xception 모델과 ASPP(Atrous Spatial Pyramid Pooling), decoder modules을 적용하여 더 빠르고 강한 encoder-decoder network를 구축했다. PASCAL VOC 2012와 Cityscapes 데이터셋에서 어떠한 post-processing 없이도 각각 89%, 82.1%의 정확도를 기록했다.
<br /><br /></p>

<h1 id="details">Details</h1>
<h2 id="도입부--관련-연구">도입부 &amp; 관련 연구</h2>
<ul>
  <li>FCN을 기반으로 한 Deep convolutional netrowrks는 benchmark tasks(BMT, 실제로 존재하는 하드웨어 혹은 소프트웨어의 성능을 비교 분석하는 방법)에서 수작업에 의존하던 시스템보다 훨씬 더 큰 개선을 이루어냈다.</li>
  <li>본 연구에서는 spatial pyramid pooling 모듈이나 encoder-decoder 구조 라는 두 가지 종류의 networks를 사용하여 semantic segmentaion task를 수행했다.</li>
  <li>PSPNet이 다양한 grid scale을 사용하여 pooling 연산을 수행할 때, 본 연구의 이전 버전의 모델인 Deeplabv3에서는 다양한 크기의 context information을 캡처하기 위해 여러 rates을 가진 atrous convolution을 병렬로 적용한 ASPP 기법을 사용했다.</li>
  <li>layer의 마지막 피처 맵에 많은 semantic information이 잘 인코딩되더라도 pooling과 stride를 통해 객체 간의 경계와 관련된 세부 정보가 누락될 수 있다. 이러한 문제는 Atrous convolution을 적용하여 더 세밀하게 피처맵을 추출함으로써 완화시킬 수 있다.</li>
  <li>하지만, 한정적인 메모리를 고려할 때 입력 이미지보다 8배, 4배 더 작은 피처맵을 추출하기에는 어려움이 있다. 더 큰 receptive field를 가지기위해(넓은 범위의 공간정보를 활용하기위해) 더 작은 피처맵을 추출한다는 것은 그만큼 더 많은 층을 거쳐야한다는 것이니까 층이 깊어질수록 훨씬 더 많은 연산량을 요구한다.</li>
  <li>본 연구에서 사용한 encoder-decoder 모델의 encoder 경로에서는 일반적인 covolution에 비해 receptive field가 커져도 atrous convolution을 통해 피처수가 증가하지 않기 때문에 빠른 연산이 가능하고, decoder 경로에서는 객체간의 경계를 점진적으로 복구시켜 위의 문제를 해결할 수 있다. 쉽게 말해, DeepLabv3와 같이 컴퓨팅 자원에 따라 atrous convolution을 적용시켜 인코더 단계에서 풍부한 semantic segmentation 정보를 저장하고, 디코더 단계에서 해상도를 높이는 기법을 사용한다.</li>
  <li>연산 속도와 정확도를 개선시킨 Xception model을 추가로 개발했고, ASPP와 encoder-decoder 구조를 사용하여 좋은 성능을 기록했다.</li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/65c1e4a4-e896-4ba9-90fc-868d5817e7e4" />
  </p>
</div>

<p><br /><br /></p>

<ul>
  <li>
    <p>DeepLabv3+ 모델은 다음과 같은 구조로 예측을 수행한다.</p>

    <p>1) Xception 혹은 ResNet을 backbone으로 사용하여 low level features(ex. ResNet에서 채널이 256인 conv2x 블록)와 GAP와 완전 결합층이 시작되기 이전의 피처(편의상 x)를 가져온다.</p>

    <p>2) Encoder 단계에서는 backbone의 x를 사용하는데 x를 가져올 때 encoder의 output stride에 맞춰서 크기가 유지되도록 하고, 논문에 근거해 특정 층(ex. ImageNet dataset이라고 가정하면 Resnet101에서 마지막 블록인conv5x의 ouput이 downsample을 진행하면 7x7이 되니까 downsample을 진행하지 않고, 오히려 atrous convolution을 사용하여 피처맵의 크기는 그대로 유지시키고 receptive field를 키움)에는 atrous convolution이 적용될 수 있다.</p>

    <p>3) Backbone에서 출력된 x에 ASPP(1x1, 3x3_rate6, 3x3_rate12, 3x3_rate18, GAP)를 수행하고 gap를 통해 작아진 output은 upsampling을 통해 크기를 다른 피처맵과 같이한 뒤 concatenate를 진행한다.</p>

    <p>4) Concatenate 후에는 1x1 합성곱을 통해 채널 수를 줄인 뒤 low_level_feature와 동일한 크기의 피처맵이 되도록 upsampling을 진행한다.</p>

    <p>5) Decoder 단계에서는 backbone에서 추출한 low level features에 1x1 합성곱을 적용해 채널은 48로 줄인뒤 이를 encoder의 출력과 다시 concatenate 한다.</p>

    <p>6) Encoder와 Decoder의 피처맵이 합쳐졌다면 3x3x256 크기를 가진 필터를 두번 거친 후 class 수만큼의 채널을 가진 1x1 합성곱을 통과한다.</p>

    <p>7) 마지막으로, 피처맵의 크기를 원본과 같은 크기로 upsampling하여 에측을 수행한다.</p>
  </li>
</ul>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/3ae598ad-7207-4357-a45e-15c783bb9dac" />
  </p>
</div>

<p><br /><br /></p>

<h2 id="methods">Methods</h2>
<h3 id="1-encoder-decoder-with-atrous-convolution">1. Encoder-Decoder with Atrous Convolution</h3>
<p><strong>Atrous convolution</strong></p>

<p>Atrous convolution은 피처의 해상도를 컨트롤 하고, multi-scale의 정보를 캡쳐하기 위해 필터의 field-of-view를 조정한다.
<br /></p>

<div align="center">
  <p>
  <img width="200" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/dd7a6934-b576-4d95-bfb5-130b80e99892" />
  </p>
</div>

<p><br /></p>

<p>여기서 r은 atrous rate으로 input signal에 얼만큼의 stride을 적용할 것인지를 의미하고(일반적인 convolution은 r=1), i는 location, k는 kernel size, w는 필터, x는 input 피처맵, y는 output 피처맵을 나타낸다. 필터의 field-of-view(receptive field)는 rate에 따라 달라진다.</p>

<p>Ex 1) rate = 1일 때, k = 3(sigma의 k는 0부터 2까지) 일 때 x는 x[0 + 1 * 0] x[0 + 1 * 1]  x[0 + 1 * 2]로 증가하기때문에</p>

<p>y[0] = x[0] * w[0] + x[1] * w[1] + x[2] * w[2]</p>

<p>Ex 2) rate = 2일 때, k = 3일 때 x는 x[0 + 2 * 0] x[0 + 2 * 1]  x[0 + 2 * 2]로 증가하기떄문에
y[0] = x[0] * w[0] + x[2] * w[1] + x[4] * w[2]</p>

<p>rate 2일때, receptive field가 더 커진 것을 알 수 있다.
<br /></p>

<p><strong>Depthwise separable convolution</strong></p>

<p>Depthwise separable convolution은 입력된 채널과 같은 채널을 갖는 필터로 각각의 필터는 자신이 담당한 입력 채널만을 상대로 연산을 수행하기때문에 필터의 모든 채널들이 입력 채널들에 대해 교차로 연산되는 일반적인 covolution보다 연산량이 훨씬 작다는 장점을 가지고 있다.</p>

<p>아래 그림은 3x3 depthwise convolution이 각각의 채널마다(separable) 하나의 필터를 배치해 나온 특징들을 결합해 pointwise convolution을 만드는데 depthwise convolution 단계에서 atrous convolution을 적용해 rate = 2로 하는 (c)와 같은 방법론을 사용한다는 예시이다. Pointwise convolution 단계에서는 1x1 convolution을 통해 depthwise convolution의 결과를 선형결합하고 채널수를 줄일 수 있다.</p>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/1fb8c1ff-886e-4630-9d86-3a3d614c8e7e" />
  </p>
</div>

<p><br /></p>

<p>DeepLabv3+ 모델에서는 depthwise separable convolution을 사용하여 일반적인 covolution을 depthwise covolution과 pointwise covolution으로 변환하여 계산 복잡도를 크게 줄였다. 특히, depthwise convolution에 atrous convolution이 적용되어 atrous separable convolution이라고도 불린다. 이는 계산 복잡도를 크게 줄이면서 성능은 유지할 수 있도록 한 좋은 방법이었다.</p>

<p><strong>DeepLabv3 as encoder</strong></p>

<p>Output stride는 최종 output resolution(Global Average Pooling이나 Fully Connected 층 이전)에 대한 입력 이미지의 spatial resolution의 비율로 나타낸다. Classsification에서는 일반적으로 output의 resolution이 input보다 32배가
작기 때문에 output stride = 32가 된다(224x224 이미지가 FC층 출력으로 사용되기 직전에 7x7로 줄어들었을 때의 예시). Segmentation task에서는 좀 더 세밀한 특징을 추출하기 위해 마지막 한개 혹은 두개의 블록을 제거한 후, atrous convolution을 적용해
output stride = 16 혹은 8로 만들어서 resolution을 높인다.</p>

<p>또한, DeepLabv3에서처럼 다양한 비율의 atrous convolution을 적용한 atrous spatial pyramid pooling module을 사용하여 multi-scale의 특징을 추출했고, logits 이전의 마지막 피처맵을 인코더의 출력으로 사용했다. 인코더의 출력은 256개의 채널로 풍부한 semantic information을 담고있다. 컴퓨팅 자원에 따라 atrous rate을 조절하면 추출되는 피처맵의 해상도를 임의로 조정할 수도 있다.</p>

<p><strong>Proposed decoder</strong></p>

<p>DeepLabv3의 인코더는 일반적으로 output stride가 16이기떄문에 피처맵은 원본 이미지보다 16배 줄어들었고 디코더 과정에서 이를 복원하기 위해 피처가 16배 upsampling 된다. 하지만, 이러한 naive decoder module은 segment의 세부 정보를 잘 복구시키지 못할 수도 있기 때문에 DeepLabv3+에서는 다른 방법을 사용한다.</p>

<p>먼저, 인코더 피처를 4배 upsampling 하고, 동일한 spatial resolution을 갖는 네트워크 백본의 저수준 피처와 결합한다. 저수준 피처가 너무 많은 채널을 가지면 합치는 과정에서 연산량이 증가할 수 있기때문에 1x1 필터를 사용해 채널 수를 줄인다. 
연결 후에는 3x3 convolution을 몇번 적용하여 특징을 업데이트 시키고 upsampling을 통해 피처맵의 크기를 4배(원본과 같이) 키운다.</p>

<p>또한, 실험을 통해 output stride가 16일 때 정확도와 연산 속도의 trade-off가 가장 좋았고, stride가 8일 때는 성능이 좀 더 좋았지만 그에 비해 연산 속도가 많이 느렸다.</p>

<h3 id="2-modified-aligned-xception">2. Modified Aligned Xception</h3>
<p>Xception 모델이 빠른 계산으로 ImageNet에서 좋은 분류 성능을 기록했고 MSRA팀이 Alined Xception 모델을 통해 객체 탐지에서 성능을 더욱 향상시켰다. 본 연구에서는 sementic segmentation 분야에서도 성능을 향상시키기 위해 다음과 같이 
MSRA 팀의 Alined Xception 모델보다 몇 가지 구조를 더 수정했다.</p>

<ul>
  <li>전체적인 구조는 크게 바꾸지 않고도 깊이를 깊게하여 연산 속도와 메모리 효율을 높였다.</li>
  <li>모든 max pooling 연산을 atrous covonlution이 적용된 depthwise separable covolution으로 바꿨다.</li>
  <li>MobileNet의 구조처럼 3x3 convolution 연산 후에는 batch normalization과 ReLU를 추가했다.</li>
</ul>

<p><br /><br /></p>

<h2 id="experimental-evaluation">Experimental Evaluation</h2>
<p>모델의 구조는 아래와 같고, 라벨은 20개의 foreground object와 1개의 backgound object 이루어져있다. 평가지표는 mIOU를 사용했고 initial lerning rate 0.007, crop size 513 x 513, output stride = 16,
학습에서는 random scale data augmentation을 사용했다.</p>

<p><br /></p>

<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/73587e4a-0e4f-4a1c-8fec-bb5480820dc3" />
  </p>
</div>

<p><br /></p>

<h3 id="1-decoder-design-choices">1. Decoder Design Choices</h3>
<p>Decoder module에서 1x1 convolution의 효과를 평가하기 위해 3x3x256의 피처맵과 Resnet-101 네트워크의 Conv2 피처(conv2x residual block의 마지막 피처맵)를 사용한다. 
먼저, 인코더 모듈에서는 저수준 피처맵의 채널을 48이나 32로 줄이면 성능이 향상되기때문에 1x1x48의 필터를 적용하고, 그 후, 디코더 모듈에는 3x3 컨불루션을 배치했다. 그 결과, 디코더 모듈의 피처맵과 인코더의 피처맵을 합쳤을 때 단순히 여러개의 convolution을 사용하는 것보다 3x3x256 합성곱을 두번 사용한 것이 더 효과적이었다. 256에서 128로 필터의 채널을 바꾸거나 필터의 크기를 3x3을 1x1으로 바꾸면 성능이 하락했다.</p>

<p>이외에도 다양한 실험을 해봤지만 결국에는 위와 같이 가장 간단하면서도 효과적인 Decoder module을 선택했다. DeepLabv3의 피처맵과 1x1 필터가 적용되어 채널이 감소된 Conv2 피처맵은 두 개의 3x3x256 필터로 연결된다. DeepLabv3+ 모델의 output stride는 최대 4까지 추천한다.</p>

<h3 id="2-resnet-101-as-network-backbone">2. ResNet-101 as Network Backbone</h3>
<p>모델의 변화에 따른 정확도와 속도를 비교하기 위해 mIOU와 Multi-Adds를 사용했다. ResNet-101을 backbone으로 사용했을 때, Atrous convolution 덕분에 하나의 model 만을 가지고도 rate을 다르게 해서 다양한 resolution을 갖는 피처를 얻어 성능을 비교해볼 수 있었다.</p>
<ul>
  <li>Baseline: output stride = 8로 하고 multi-scale input을 사용하면 성능이 향상된다. 좌우반전을 적용한 이미지를 추가하면 계산 복잡도는 두배가 되지만 그에 비해 성능 향상은 미비했다.</li>
  <li>Adding decoder: ouptut stride를 8로 하면 16으로 설정했을 때보다 성능이 더 좋지만 약 20억개에 달하는 오버헤드가 발생한다. Multi-scale과 좌우반전을 적용하면 성능이 향상된다.</li>
  <li>Coarser feature maps: atrous convolution을 사용하지 않고 빠른 연산을 위해 train output stride를 32, 디코더를 추가했을 때 성능이 2% 개선됐지만 train output stride = 16, eval output stride를 여러 값으로 사용했을 때보다
성능이 떨어졌다. 따라서, train, test output stride는 16이나 8을 쓸 것을 권장한다.</li>
</ul>

<h3 id="3-xception-as-network-backbone">3. Xception as Network Backbone</h3>
<ul>
  <li>ImangeNet pretraining: 큰 하이퍼파라미터 조정 없이 ImageNet pretraining을 시켰을 때 배치 정규화나 ReLU를 추가하지 않은 Alined Xception에서 성능이 하락했다.</li>
  <li>Baseline: 디코더를 사용하고, output stride = 16으로 통일시켰을 때 Xception 모델은 ResNet-101보다 성능이 약 2% 높았다. Test output stride = 8 로 하고 좌우반전을 적용하면 추가적인 성능 향상을 얻을 수 있다.</li>
  <li>Adding decoder: 디코더를 출력하면 eval test = 16으로 사용했을 때 약 0.8% 성능이 올라간다.</li>
  <li>Using depthwise separable covolution: depthwise separable convolution을 ASPP와 디코더 모듈에 추가했을 때 연산 복잡도는 크게 감소했는데도 비슷한 mIOU score를 얻었다.</li>
  <li>Pretraining on COCO: MS-COCO dataset을 추가로 사전학습 하여 성능을 2% 개선시켰다.</li>
  <li>Pretraining on JFT: JFT-300M dataset을 추가로 사전학습 하여 성능을 1% 개선시켰다.</li>
  <li>Test set resluts: 계산 복잡도를 고려하지 않아서 output stride = 8, 배치 정규화를 사용하여 학습했을 때 JFT dataset의 사전학습과는 관계없이 mIOU 87.8%와 89%의 성능을 달성했다.</li>
  <li>Failure mode: DeepLabv3+ 모델은 소파와 의자, 많이 가려진 물체, 잘 보이지 않는 물체를 잘 세분화하지 못한다.</li>
</ul>

<h3 id="4-improvement-along-object-boundaries">4. Improvement along Object Boundaries</h3>
<p><br /></p>
<div align="center">
  <p>
  <img width="700" alt="image" src="https://github.com/Hyeonseung0103/Hyeonseung0103.github.io/assets/97672187/854ac372-d2fe-4440-a96b-a346164cc4e9" />
  </p>
</div>

<p><br /></p>

<p>Object의 Boundaries 근처에서 proposed decoder를 활용하여 정확도를 높이기 위해 trimap experiment를 수행했다. Trimap은 배경, 전경, 배경과 전경 사이에 있는 불확실한 영역을 분류하는 작업을 수행하는 과정이라고 할 수 있다.
일반적으로 객체 경계 주변에서 발생하는 val set의 비어있는 label annotations(불확실한 영역)에 morphological dilation(trimap width를 조정해서 해당 영역을 더 부각시킴)을 적용했다.</p>

<p><br /></p>

<p>그 후, dilated band(확장된 trimap)에 대해 mIOU를 구했다. 단순히 naive decoder를 사용했을 때보다 proposed decoder를 사용했을 때 ResNet과 Xception 모델에서 4.8%, 5.4%의 mIOU 향상을 기록했고 dilated band가 좁을수록 성능은 더 잘 향상되었다.</p>

<h3 id="5-experimental-results-on-cityscapes">5. Experimental Results on Cityscapes</h3>
<p>Cityscapes dataset에서도 DeepLabv3+를 실험해봤다. Aligned Xception 모델을 backbone으로 사용하고 DeepLabv3의 ASPP를 사용했을 때 proposed decoder를 적용하면 mIOU가 78.79%로 naive decoder를 사용했을 때보다 약 1.46% 성능이 향상됐다.
Augmentation image를 제거했을 때 성능이 79.14%로 향상된 것을 보아 DeepLab 모델에서는 PASCAL VOC 2012 dataset에서 augmentation이 효과적이고 Cityscapes에서는 그렇지 않다는 것을 알 수 있다.</p>

<p>DeepLabv3+ 모델은 네트워크의 layers를 더 추가한 결과(X-71) val set에서 79.55%, test set에서는 82.1%의 최고 성능을 달성했다.</p>

<h2 id="conclusion">Conclusion</h2>
<ul>
  <li>DeepLabv3+ 모델은 풍부한 cotext information을 위한 DeepLabv3의 인코딩과 Object Boundaries를 잘 복구하기 위한 디코더 모듈로 이루어진 encoder-decoder 구조를 사용한다.</li>
  <li>컴퓨팅 자원에 따라 atrous convolution을 적용하여 인코더의 특징을 임의의 resolution으로 추출할 수도 있다.</li>
  <li>Xception 모델과 atrous separable convolution을 사용하여 더 빠르고 좋은 모델을 만들었다.</li>
  <li>PASCAL VOC 2012와 Cityscapes datasets에서 SOTA를 기록했다.</li>
</ul>

<p><br /><br /></p>

<h1 id="개인적인-생각">개인적인 생각</h1>
<ul>
  <li>DeepLabv1 이후 지속적으로 연구하여 architecture에 큰 변화없이 모델을 발전시켜서 SOTA를 기록해나갔다는 점이 인상깊었다.</li>
  <li>DeepLabv3+도 U-Net처럼 인코더-디코더 구조를 사용했는데, Segmentation 모델에 인코더-디코더 구조를 사용하지않으면서 성능을 좋게하는 방법은 없을지 궁금하다.</li>
  <li>ResNet101 모델을 backbone으로 사용할 때 마지막 residual block 뿐만 아니라 처음부터 atrous convolution을 사용했다면 성능이 어떻게 변화했을지 궁금하다.</li>
  <li>Pooling을 convolution 연산으로 대체하고, atrous convolution을 사용하고, fully connected layer를 fully convolutional layer로 대체하며 공간적인 정보를 최대한 보존하려고하는데 향후 논문에서는 이 공간 정보를 어떻게 더 잘 보존하는지 기대된다.</li>
</ul>

<p><br /><br /></p>

<h1 id="구현">구현</h1>
<p>DeepLabv3+를 pytorch로 구현해보자.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="n">nn</span>
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="n">F</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">conv3x3</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">1</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">stride</span><span class="o">=</span><span class="n">stride</span><span class="p">,</span> <span class="n">padding</span> <span class="o">=</span> <span class="n">dilation</span><span class="p">,</span> <span class="n">dilation</span><span class="o">=</span><span class="n">dilation</span><span class="p">,</span> <span class="n">bias</span> <span class="o">=</span> <span class="bp">False</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">conv1x1</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">stride</span><span class="o">=</span><span class="n">stride</span><span class="p">,</span> <span class="n">bias</span> <span class="o">=</span> <span class="bp">False</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">Bottleneck</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="n">expansion</span> <span class="o">=</span> <span class="mi">4</span>
    
    <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">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">downsample</span> <span class="o">=</span> <span class="bp">None</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">Bottleneck</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">conv1</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">bn1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">conv2</span> <span class="o">=</span> <span class="n">conv3x3</span><span class="p">(</span><span class="n">out_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span><span class="p">,</span> <span class="n">dilation</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">bn2</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">conv3</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="n">out_channels</span><span class="p">,</span> <span class="n">out_channels</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">expansion</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">bn3</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span> <span class="o">*</span> <span class="bp">self</span><span class="p">.</span><span class="n">expansion</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">relu</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">ReLU</span><span class="p">(</span><span class="n">inplace</span> <span class="o">=</span> <span class="bp">True</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">downsample</span> <span class="o">=</span> <span class="n">downsample</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">stride</span> <span class="o">=</span> <span class="n">stride</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="n">identity</span> <span class="o">=</span> <span class="n">x</span>

        <span class="c1"># 1x1 필터
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">conv1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">bn1</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">relu</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

        <span class="c1"># 3x3 필터
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">conv2</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">bn2</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">relu</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

        <span class="c1"># 1x1 필터
</span>        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">conv3</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">bn3</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

        <span class="k">if</span> <span class="bp">self</span><span class="p">.</span><span class="n">downsample</span> <span class="ow">is</span> <span class="ow">not</span> <span class="bp">None</span><span class="p">:</span>
            <span class="n">identity</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">downsample</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        
        <span class="n">out</span> <span class="o">+=</span> <span class="n">identity</span>
        <span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">relu</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>

        <span class="k">return</span> <span class="n">out</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">ResNet</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">block</span><span class="p">,</span> <span class="n">layers_num</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">ResNet</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        
        <span class="n">multi_grid</span> <span class="o">=</span> <span class="p">[</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">4</span><span class="p">]</span>
        <span class="s">'''
        이 multi_grid 비율에 맞게 rate이 곱해져서 atrous convolution
        이전 버전인 deeplabv3 논문에 근거. For example, when output stride = 16 and Multi Grid = (1, 2, 4), 
        the three convolutions will have rates = 2 · (1, 2, 4) = (2, 4, 8) in the block4, respectively.
        '''</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span> <span class="o">=</span> <span class="mi">64</span>
        
        <span class="bp">self</span><span class="p">.</span><span class="n">conv1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">(</span><span class="mi">3</span><span class="p">,</span> <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">7</span><span class="p">,</span> <span class="n">stride</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">3</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">bn1</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">relu</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">ReLU</span><span class="p">(</span><span class="n">inplace</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">maxpool</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">MaxPool2d</span><span class="p">(</span><span class="n">kernel_size</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">stride</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">padding</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">layer1</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_make_layer</span><span class="p">(</span><span class="n">block</span><span class="p">,</span> <span class="mi">64</span><span class="p">,</span> <span class="n">layers_num</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">layer2</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_make_layer</span><span class="p">(</span><span class="n">block</span><span class="p">,</span> <span class="mi">128</span><span class="p">,</span> <span class="n">layers_num</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">layer3</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_make_layer</span><span class="p">(</span><span class="n">block</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="n">layers_num</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">2</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">layer4</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">_make_atrous_layer</span><span class="p">(</span><span class="n">block</span><span class="p">,</span> <span class="mi">512</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="n">multi_grid</span> <span class="o">=</span> <span class="n">multi_grid</span><span class="p">)</span>
        <span class="c1"># output stride = 16이면 최종 출력값이 14라서 downsampling이 되면 안 됨. 오히려, atrous convolution으로 receptive field를 넓혔음.
</span>
        <span class="c1"># Deeplab에서는 residual block 까지만 사용.
</span>        <span class="c1"># self.avgpool = nn.AdaptiveAvgPool2d((1,1))
</span>        <span class="c1"># self.fc = nn.Linear(512 * block.expansion, num_classes)
</span>
        <span class="c1"># 가중치 초기화
</span>        <span class="k">for</span> <span class="n">m</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">modules</span><span class="p">():</span>
            <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">):</span>
                <span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">kaiming_normal_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">,</span> <span class="n">mode</span> <span class="o">=</span> <span class="s">'fan_out'</span><span class="p">,</span> <span class="n">nonlinearity</span><span class="o">=</span><span class="s">'relu'</span><span class="p">)</span>
            <span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">):</span>
                <span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">constant_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>
                <span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">constant_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">bias</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>

    
    <span class="k">def</span> <span class="nf">_make_layer</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">block</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">blocks_num</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">):</span>
        <span class="n">downsample</span> <span class="o">=</span> <span class="bp">None</span>

        <span class="k">if</span> <span class="n">stride</span> <span class="o">!=</span> <span class="mi">1</span> <span class="ow">or</span> <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span> <span class="o">!=</span> <span class="n">out_channels</span> <span class="o">*</span> <span class="n">block</span><span class="p">.</span><span class="n">expansion</span><span class="p">:</span>
            <span class="n">downsample</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span>
                <span class="n">conv1x1</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span> <span class="o">*</span> <span class="n">block</span><span class="p">.</span><span class="n">expansion</span><span class="p">,</span> <span class="n">stride</span><span class="p">),</span>
                <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span> <span class="o">*</span> <span class="n">block</span><span class="p">.</span><span class="n">expansion</span><span class="p">)</span>
            <span class="p">)</span>
        
        <span class="n">layers</span> <span class="o">=</span> <span class="p">[]</span>
        
        <span class="c1"># conv2,3,4x에서는 dilation이 항상 1이니까 굳이 block을 만들때 rate을 지정해줄 필요 없음.
</span>        <span class="n">layers</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">block</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span><span class="p">,</span> <span class="n">downsample</span><span class="p">))</span>  
        <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span> <span class="o">=</span> <span class="n">out_channels</span> <span class="o">*</span> <span class="n">block</span><span class="p">.</span><span class="n">expansion</span> 
        <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">blocks_num</span><span class="p">):</span>
            <span class="n">layers</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">block</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">))</span>
        <span class="k">return</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span><span class="o">*</span><span class="n">layers</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">_make_atrous_layer</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">block</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">multi_grid</span> <span class="o">=</span> <span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">4</span><span class="p">]):</span>
        <span class="c1"># resnet 가장 마지막 층에 사용하는 atrous_convolution에서는 output stride를 유지하기위해 downsampling이 일어나지않음.
</span>        <span class="n">layers</span> <span class="o">=</span> <span class="p">[]</span>
        <span class="n">layers</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">block</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">stride</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="n">multi_grid</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">*</span> <span class="n">dilation</span><span class="p">))</span> 
        <span class="c1"># 첫번째 residual block에 dilation 2 적용.
</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span> <span class="o">=</span> <span class="n">out_channels</span> <span class="o">*</span> <span class="n">block</span><span class="p">.</span><span class="n">expansion</span>
        <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="mi">1</span><span class="p">,</span> <span class="nb">len</span><span class="p">(</span><span class="n">multi_grid</span><span class="p">)):</span>
            <span class="n">layers</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">block</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="n">multi_grid</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">*</span> <span class="n">dilation</span><span class="p">))</span>

        <span class="k">return</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span><span class="o">*</span><span class="n">layers</span><span class="p">)</span>
        
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">conv1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">bn1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">relu</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">maxpool</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>

        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">layer1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">low_level_features</span> <span class="o">=</span> <span class="n">x</span>

        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">layer2</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">layer3</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">layer4</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        
        <span class="c1"># x = self.avgpool(x)
</span>        <span class="c1"># x = torch.flatten(x, 1)
</span>        <span class="c1"># x = self.fc(x)
</span>
        <span class="k">return</span> <span class="n">x</span><span class="p">,</span> <span class="n">low_level_features</span>

    <span class="k">def</span> <span class="nf">_init_weight</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">m</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">modules</span><span class="p">():</span>
            <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">):</span>
                <span class="n">torch</span><span class="p">.</span><span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">kaiming_normal_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">)</span>
            <span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">):</span>
                <span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">fill_</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
                <span class="n">m</span><span class="p">.</span><span class="n">bias</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">zero_</span><span class="p">()</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">ASPP</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">kernel_size</span><span class="p">,</span> <span class="n">padding</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">1</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">ASPP</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">aspp</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">Sequential</span><span class="p">(</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">(</span><span class="n">in_channels</span><span class="p">,</span> <span class="n">out_channels</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="n">kernel_size</span><span class="p">,</span> <span class="n">stride</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> 
                      <span class="n">padding</span><span class="o">=</span><span class="n">padding</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="n">dilation</span><span class="p">,</span> <span class="n">bias</span> <span class="o">=</span> <span class="bp">False</span><span class="p">),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">(</span><span class="n">out_channels</span><span class="p">),</span>
            <span class="n">nn</span><span class="p">.</span><span class="n">ReLU</span><span class="p">()</span>
        <span class="p">)</span>

    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">aspp</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">_init_weight</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">m</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">modules</span><span class="p">():</span>
            <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">):</span>
                <span class="n">torch</span><span class="p">.</span><span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">kaiming_normal_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">)</span>
            <span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">):</span>
                <span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">fill_</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
                <span class="n">m</span><span class="p">.</span><span class="n">bias</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">zero_</span><span class="p">()</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># output stride가 16일 때
# rates = 1, 6, 12, 18
</span><span class="k">class</span> <span class="nc">DeepLabv3Plus</span><span class="p">(</span><span class="n">nn</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <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">num_classes</span><span class="p">):</span>
        <span class="nb">super</span><span class="p">(</span><span class="n">DeepLabv3Plus</span><span class="p">,</span> <span class="bp">self</span><span class="p">).</span><span class="n">__init__</span><span class="p">()</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">resnet</span> <span class="o">=</span> <span class="n">ResNet</span><span class="p">(</span><span class="n">Bottleneck</span><span class="p">,</span> <span class="p">[</span><span class="mi">3</span><span class="p">,</span> <span class="mi">4</span><span class="p">,</span> <span class="mi">23</span><span class="p">,</span> <span class="mi">3</span><span class="p">])</span>

        
        <span class="c1"># deeplabv3 논문. atrous convolution -&gt; "all with 256 filters and batch normalization"
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">aspp1</span> <span class="o">=</span> <span class="n">ASPP</span><span class="p">(</span><span class="mi">2048</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span> <span class="n">dilation</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">aspp2</span> <span class="o">=</span> <span class="n">ASPP</span><span class="p">(</span><span class="mi">2048</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">6</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">6</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">aspp3</span> <span class="o">=</span> <span class="n">ASPP</span><span class="p">(</span><span class="mi">2048</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">12</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">12</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">aspp4</span> <span class="o">=</span> <span class="n">ASPP</span><span class="p">(</span><span class="mi">2048</span><span class="p">,</span> <span class="mi">256</span><span class="p">,</span> <span class="n">kernel_size</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">18</span><span class="p">,</span> <span class="n">dilation</span> <span class="o">=</span> <span class="mi">18</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">avgpool</span> <span class="o">=</span> <span class="n">nn</span><span class="p">.</span><span class="n">AdaptiveAvgPool2d</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">pool_conv</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="mi">2048</span><span class="p">,</span> <span class="mi">256</span><span class="p">)</span> <span class="c1"># 풀링 후 풀링 채널도 256으로 통일
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">encoder_conv1x1</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="mi">256</span> <span class="o">*</span> <span class="mi">5</span><span class="p">,</span> <span class="mi">256</span><span class="p">)</span>

        <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv1x1</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="mi">256</span><span class="p">,</span> <span class="mi">48</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv3x3_a</span> <span class="o">=</span> <span class="n">conv3x3</span><span class="p">(</span><span class="mi">256</span> <span class="o">+</span> <span class="mi">48</span><span class="p">,</span> <span class="mi">256</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv3x3_b</span> <span class="o">=</span> <span class="n">conv3x3</span><span class="p">(</span><span class="mi">256</span><span class="p">,</span> <span class="mi">256</span><span class="p">)</span>
        <span class="bp">self</span><span class="p">.</span><span class="n">last_conv</span> <span class="o">=</span> <span class="n">conv1x1</span><span class="p">(</span><span class="mi">256</span><span class="p">,</span> <span class="n">num_classes</span><span class="p">)</span>
    
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="nb">input</span><span class="p">):</span>
        <span class="c1"># 인코더
</span>        <span class="n">x</span><span class="p">,</span> <span class="n">low_level_features</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">resnet</span><span class="p">(</span><span class="nb">input</span><span class="p">)</span>
        <span class="n">x1</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">aspp1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x2</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">aspp2</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x3</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">aspp3</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x4</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">aspp4</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x5</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">avgpool</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x5</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">pool_conv</span><span class="p">(</span><span class="n">x5</span><span class="p">)</span>
        <span class="n">x5</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">upsample</span><span class="p">(</span><span class="n">x5</span><span class="p">,</span> <span class="n">size</span> <span class="o">=</span> <span class="n">x4</span><span class="p">.</span><span class="n">size</span><span class="p">()[</span><span class="mi">2</span><span class="p">:],</span> <span class="n">mode</span><span class="o">=</span><span class="s">'bilinear'</span><span class="p">)</span> <span class="c1"># pooling으로 1,1로 줄어들었으니까 다시 다른 피처맵과 같은 크기로
</span>        
        <span class="n">x</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">x1</span><span class="p">,</span><span class="n">x2</span><span class="p">,</span><span class="n">x3</span><span class="p">,</span><span class="n">x4</span><span class="p">,</span><span class="n">x5</span><span class="p">],</span> <span class="mi">1</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">encoder_conv1x1</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
        <span class="n">x</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">upsample</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">size</span> <span class="o">=</span> <span class="p">(</span><span class="nb">input</span><span class="p">.</span><span class="n">size</span><span class="p">()[:</span><span class="o">-</span><span class="mi">2</span><span class="p">]</span> <span class="o">//</span> <span class="mi">4</span><span class="p">,</span> <span class="nb">input</span><span class="p">.</span><span class="n">size</span><span class="p">()[:</span><span class="o">-</span><span class="mi">1</span><span class="p">]</span> <span class="o">//</span> <span class="mi">4</span><span class="p">),</span> <span class="n">mode</span> <span class="o">=</span> <span class="s">'bilinear'</span><span class="p">)</span>
        <span class="c1"># decoder와 합치기위해 low_level_feature와 같은 크기로 low level feature는 ResNet에서 input보다 4배 작은 크기로 넘어옴
</span>
        <span class="c1"># 디코더
</span>        <span class="n">xd</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv1x1</span><span class="p">(</span><span class="n">low_level_features</span><span class="p">)</span>
        <span class="n">xd</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="n">cat</span><span class="p">([</span><span class="n">x</span><span class="p">,</span> <span class="n">xd</span><span class="p">],</span> <span class="mi">1</span><span class="p">)</span> 
        <span class="n">xd</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv3x3_a</span><span class="p">(</span><span class="n">xd</span><span class="p">)</span>
        <span class="n">xd</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">decoder_conv3x3_b</span><span class="p">(</span><span class="n">xd</span><span class="p">)</span>
        <span class="n">xd</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">last_conv</span><span class="p">(</span><span class="n">xd</span><span class="p">)</span>
        <span class="n">output</span> <span class="o">=</span> <span class="n">F</span><span class="p">.</span><span class="n">upsample</span><span class="p">(</span><span class="n">xd</span><span class="p">,</span> <span class="n">size</span> <span class="o">=</span> <span class="nb">input</span><span class="p">.</span><span class="n">size</span><span class="p">()[</span><span class="mi">2</span><span class="p">:],</span> <span class="n">mode</span> <span class="o">=</span> <span class="s">'bilinear'</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">output</span>

    <span class="k">def</span> <span class="nf">_init_weight</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">m</span> <span class="ow">in</span> <span class="bp">self</span><span class="p">.</span><span class="n">modules</span><span class="p">():</span>
            <span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">Conv2d</span><span class="p">):</span>
                <span class="n">torch</span><span class="p">.</span><span class="n">nn</span><span class="p">.</span><span class="n">init</span><span class="p">.</span><span class="n">kaiming_normal_</span><span class="p">(</span><span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">)</span>
            <span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">nn</span><span class="p">.</span><span class="n">BatchNorm2d</span><span class="p">):</span>
                <span class="n">m</span><span class="p">.</span><span class="n">weight</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">fill_</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
                <span class="n">m</span><span class="p">.</span><span class="n">bias</span><span class="p">.</span><span class="n">data</span><span class="p">.</span><span class="n">zero_</span><span class="p">()</span>
</code></pre></div></div>

<h1 id="이미지-출처">이미지 출처</h1>
<ul>
  <li><a href="https://browse.arxiv.org/pdf/1802.02611.pdf">Encoder-Decoder with Atrous Separable Convolution for Semantic Image Segmentation Paper</a></li>
  <li><a href="https://arxiv.org/pdf/1606.00915v2.pdf">Atrous convolution DeepLabv2 논문 참고</a></li>
  <li><a href="https://medium.com/mini-distill/pps-deeplab-semantic-image-segmentation-with-deep-convolutional-nets-atrous-convolution-and-ae6463029410">Atrous convolution 이미지 예시</a></li>
  <li><a href="https://itrepo.tistory.com/32">Receptive field 이미지</a></li>
  <li><a href="https://devocean.sk.com/blog/techBoardDetail.do?ID=163989">Trimap 예시</a></li>
</ul>]]></content><author><name>{&quot;name&quot;=&gt;nil, &quot;avatar&quot;=&gt;nil, &quot;bio&quot;=&gt;nil, &quot;location&quot;=&gt;nil, &quot;email&quot;=&gt;nil, &quot;links&quot;=&gt;[{&quot;label&quot;=&gt;&quot;Email&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-envelope-square&quot;}, {&quot;label&quot;=&gt;&quot;Website&quot;, &quot;icon&quot;=&gt;&quot;fas fa-fw fa-link&quot;}, {&quot;label&quot;=&gt;&quot;Twitter&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-twitter-square&quot;}, {&quot;label&quot;=&gt;&quot;Facebook&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-facebook-square&quot;}, {&quot;label&quot;=&gt;&quot;GitHub&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-github&quot;}, {&quot;label&quot;=&gt;&quot;Instagram&quot;, &quot;icon&quot;=&gt;&quot;fab fa-fw fa-instagram&quot;}]}</name></author><category term="Segmentation" /><summary type="html"><![CDATA[Atrous Convolution이란? DeepLabv3+ 논문을 리뷰하기 앞서 DeepLab 모델들에서 사용된 핵심 기법인 Atrous Convolution에 대해 알아보자.]]></summary></entry></feed>