위에 적힌 두 Objective가 실제로 사용되는 metric이 아닌 surrogate function이고 내가 정말 필요로 하는 목적과는 다르다.
(sorting, top-k search)에 end-to-end tailored되어 있지 않다는 문제가 있다.
정확히 매칭되는 그림은 아닌데, 아무튼 이런 현상이 존재한다. Epoch 100 이후에 Loss는 계속 줄고 있지만, Accuracy는 더 이상 거의 증가하지 않는다.
이 논문에서는 미분 가능한 soft-정렬 연산자를 제시함으로 이런 문제를 해결하고, End-to-end style로 Sort/Permutation/TopK를 optimize하기 위한 미분 가능한 Operator를 제시했다.
Objective involving permutation
Permutation
Permutation π란 [1, 3, 2]와 같이 1부터 n이 하나씩 유일하게 들어있는 vector를 말한다.
π=[3,2,1]일 때, 다른 array s[a,b,c]에 π를 적용하면, [c,b,a]가 된다. 보통 python에서 자주 쓰는 notaiton으로 적는다면, [c,b,a]=s[π]가 되고, s에 대한 내림차순 정렬이 된다. 앞으로 우리가 관심있는 Permutation은 대체로 정렬이므로, 편의상 z를 s에 대한 정렬 permutation, z=sort(s)이라 하자.
Represent Permuation as Matrix
이런 permutation은 또한 matrix로 나타낼 수 있다.
P(z)s=001010100[a,b,c]T=[c,b,a]T
P(z)[k,:]는 j'th element가 1이고, 나머지는 다 0이다. s에서 k번째로 큰 element는 sj라는 의미이다.
이런 정렬을 하는 matrix를 Pz라 하자.
P는 이렇게 계산할 수 있다.
P[i,j]={10ifj=argmax+1−2i)s−A1otherwise
where A[i,j]=∣si−sj∣.
불연속적인 P(z)의 row를 relax하는 식으로 Permutation 혹은 Sort를 미분 가능하게 바꾸는 간단한 방법을 제안했다. 그냥 argmax를 softmax with temperature τ로 바꿔준 꼴이고, 엄청 간단하다.
P^[i,:]={10ifj=softmax[τ−1[(n+1−2i)s−A1]]otherwise where A[i,j]=∣si−sj∣
이러한 미분 불가능한 연산(정렬)을 미분 가능한 연산으로 relax하는 일반적인 방법은 이를 continuous한 공간에 매핑하는 것이다. 비슷한 예시로는, 흔히 사용되는 Classification objective가 있다.
0/1 loss for binary classification -> logistic or hinge loss
thresholding gives a mapping from real value to discrete decision.
이 matrix의 다음과 같은 성질을 보존하는 Real-value matrix를 제안햇다.
P^는 다음과 같은 성질을 만족하는 행렬이다. (P도 이를 만족한다)
Non-negativity: U[i,j]≥0∀i,j∈{1,2,...,n}
Row Affinity: ∑j=1nU[i,j]=1
Argmax Permuatiotn: u=[u1,u2,...,un] where ui=argmaxj(U[i,:]) then u is a valid permutation of {n}
(1), (2)는 discrete permutation에 대한 stochastic relaxation을 의미한다.
P[k,:]는 j'th elementh가 1이고, 나머지는 다 0이다. 즉, s에서 k번째로 큰 element는 sj라는 의미였는데, P^[k,:]는 k번째로 큰 element가 무엇인지에 대한 확률분포(에 가까운 것)이라 생각할 수 있다. P^는 argmax로 정의되어 있던 P의 i번째 row를 Softmax로 대체한 것이다.
u=[u1,...,un], where ui=argmax(P[i,:]) , then u = sort(s)
τ→+0인 경우 P^→P이다.
P^은 미분 가능한 정렬 연산자이다.
Gumbel Trick을 이용한 Stochastic Permuation and model update
Px: vector x에 대한 permutation matrix.
P^x: vector x에 대한 relaxed permuation matrix.
Gumbel Distribution
g=gumbel(0,1)=−log(−log(uniform(0,1))
이다. 길이가 n인 vector v가 있고, n개의 i.i.d 한 g1,g2,...,gn이 있다고 하자.
I1=argmax([v1+g1,v2+g2,...,vn+gn])$이면,
Pr(I1=i)=∑jexp(vj)exp(vi)
이 성립한다. 이는 조금 더 일반화될 수 있다.
k 이 n보다 작거나 같을 때, I1,...,Ik=argtop-k([v1+g1,v2+g2,...,vn+gn])이라면,
Pr(I1=i1,I2=i2,...,Ik=ik)=∏i=1k∑j∈/(I1,...,Ik−1)exp(vj)exp(vi)
가 성립한다.
v=log(s)로 생각하고, v+g에 대한 Relaxed Permutation matrix를 여러 개 만들어서 Objective를 업데이트해도 괜찮은 것이다.
∇sL=Eg[∇sf(P^sort(logs+z))]이 성립한다. (By REINFORCE)
읽고 든 생각
Neuralsort Github에 구현체가 있는데, 코드가 그렇게 어렵지 않아서, 읽어보면 재밌고, 내가 하는 연구에 어떻게던 써먹을 여지가 있지 않을까 하는 생각이 들었다.
내가 아는 한에서 미분 불가능한 Objective(주로 combinatorial problem)을 Neural Network으로 푸는 방법은 REINFORCE를 이용하는 거였다.
∇θL=Ez∼q[f(z)∇θlog(q(z∥s;θ)]
이 식을
∇βL=Ez∼q[f(z;β)∇βlog(q(z∥s;θ)]+Ez∼q[∇βf(z;β)]로 생각해서, β에 대해 optimize를 하는 방식이라 되게 신기했다. 비슷한 류의 work이 몇개 더 있는데, Gumbel-Trick을 이용해 without replacement의 Categorical Distribution을 sampling하는 방법도 있고...
Comments