1 #ifndef REDUKT_EARTHMOVERSDISTANCE_H
2 #define REDUKT_EARTHMOVERSDISTANCE_H
17 virtual float distance(
const F& feature1,
const F& feature2 )
const =0;
25 #define EMD_INFINITY 1e20
68 template<
class T,
class F>
77 float value()
const {
return m_emd; }
90 enum { MAX_SIG_SIZE =100,
106 struct node2_t *NextC;
107 struct node2_t *NextR;
112 float _C[MAX_SIG_SIZE1][MAX_SIG_SIZE1];
113 node2_t _X[MAX_SIG_SIZE1*2];
116 node2_t *_EndX, *_EnterX;
117 char _IsX[MAX_SIG_SIZE1][MAX_SIG_SIZE1];
118 node2_t *_RowsX[MAX_SIG_SIZE1], *_ColsX[MAX_SIG_SIZE1];
123 float init(
const T& Signature1,
const T& Signature2,
const GroundDistance<F>& dist );
124 void findBasicVariables(node1_t *U, node1_t *V);
125 int isOptimal(node1_t *U, node1_t *V);
126 int findLoop(node2_t **Loop);
128 void russel(
double *S,
double *D);
129 void addBasicVariable(
int minI,
int minJ,
double *S,
double *D,
130 node1_t *PrevUMinI, node1_t *PrevVMinJ,
134 void printSolution();
170 template<
class T,
class F>
179 node1_t U[MAX_SIG_SIZE1], V[MAX_SIG_SIZE1];
184 int* FlowSize = NULL;
187 w = init(Signature1, Signature2, Dist);
189 if( DEBUG_LEVEL > 1 ) {
190 printf(
"\nINITIAL SOLUTION:\n");
194 if (_n1 > 1 && _n2 > 1)
196 for (itr = 1; itr < MAX_ITERATIONS; itr++)
199 findBasicVariables(U, V);
208 if( DEBUG_LEVEL > 1 ) {
209 printf(
"\nITERATION # %d \n", itr);
214 if (itr == MAX_ITERATIONS)
215 fprintf(stderr,
"emd: Maximum number of iterations has been reached (%d)\n",
223 for(XP=_X; XP < _EndX; XP++)
227 if (XP->i == Signature1.numFeatures() || XP->j == Signature2.numFeatures())
233 totalCost += (double)XP->val * _C[XP->i][XP->j];
238 FlowP->amount = XP->val;
243 *FlowSize = FlowP-Flow;
246 printf(
"\n*** OPTIMAL SOLUTION (%d ITERATIONS): %f ***\n", itr, totalCost);
251 m_emd = (float)(totalCost / w );
258 template<
class T,
class F>
263 double sSum, dSum, diff;
264 double S[MAX_SIG_SIZE1], D[MAX_SIG_SIZE1];
266 _n1 = Signature1.numFeatures();
267 _n2 = Signature2.numFeatures();
269 if (_n1 > MAX_SIG_SIZE || _n2 > MAX_SIG_SIZE)
271 fprintf(stderr,
"emd: Signature size is limited to %d\n", MAX_SIG_SIZE);
277 for( i=0; i < _n1; ++i )
278 for( j=0; j < _n2; ++j )
280 _C[i][j] = Dist.
distance( Signature1.feature(i), Signature2.feature(j) );
281 if( _C[i][j] > _maxC )
287 for(i=0; i < _n1; i++)
289 S[i] = Signature1.weight(i);
290 sSum += Signature1.weight(i);
294 for(j=0; j < _n2; j++)
296 D[j] = Signature2.weight(j);
297 dSum += Signature2.weight(j);
303 if (fabs(diff) >=
EPSILON * sSum)
307 for (j=0; j < _n2; j++)
315 for (i=0; i < _n1; i++)
324 for (i=0; i < _n1; i++)
325 for (j=0; j < _n2; j++)
329 _maxW = sSum > dSum ? sSum : dSum;
336 return sSum > dSum ? dSum : sSum;
344 template<
class T,
class F>
348 int UfoundNum, VfoundNum;
349 node1_t u0Head, u1Head, *CurU, *PrevU;
350 node1_t v0Head, v1Head, *CurV, *PrevV;
353 u0Head.Next = CurU = U;
354 for (i=0; i < _n1; i++)
360 (--CurU)->Next = NULL;
364 v0Head.Next = _n2 > 1 ? V+1 : NULL;
365 for (j=1; j < _n2; j++)
371 (--CurV)->Next = NULL;
379 v1Head.Next->Next = NULL;
382 UfoundNum=VfoundNum=0;
383 while (UfoundNum < _n1 || VfoundNum < _n2)
386 if( DEBUG_LEVEL > 3 ) {
387 printf(
"UfoundNum=%d/%d,VfoundNum=%d/%d\n",UfoundNum,_n1,VfoundNum,_n2);
389 for(CurU = u0Head.Next; CurU != NULL; CurU = CurU->Next)
390 printf(
"[%d]",CurU-U);
393 for(CurU = u1Head.Next; CurU != NULL; CurU = CurU->Next)
394 printf(
"[%d]",CurU-U);
397 for(CurV = v0Head.Next; CurV != NULL; CurV = CurV->Next)
398 printf(
"[%d]",CurV-V);
401 for(CurV = v1Head.Next; CurV != NULL; CurV = CurV->Next)
402 printf(
"[%d]",CurV-V);
411 for (CurV=v1Head.Next; CurV != NULL; CurV=CurV->Next)
416 for (CurU=u0Head.Next; CurU != NULL; CurU=CurU->Next)
422 CurU->val = _C[i][j] - CurV->val;
424 PrevU->Next = CurU->Next;
425 CurU->Next = u1Head.Next != NULL ? u1Head.Next : NULL;
432 PrevV->Next = CurV->Next;
441 for (CurU=u1Head.Next; CurU != NULL; CurU=CurU->Next)
446 for (CurV=v0Head.Next; CurV != NULL; CurV=CurV->Next)
452 CurV->val = _C[i][j] - CurU->val;
454 PrevV->Next = CurV->Next;
455 CurV->Next = v1Head.Next != NULL ? v1Head.Next: NULL;
462 PrevU->Next = CurU->Next;
469 fprintf(stderr,
"emd: Unexpected error in findBasicVariables!\n");
470 fprintf(stderr,
"This typically happens when the EPSILON defined in\n");
471 fprintf(stderr,
"emd.h is not right for the scale of the problem.\n");
473 throw(
"EarthMoversDistance: Unexpected error in findBasicVariables!");
483 template<
class T,
class F>
486 double delta, deltaMin;
487 int i, j, minI, minJ;
491 for(i=0; i < _n1; i++)
492 for(j=0; j < _n2; j++)
495 delta = _C[i][j] - U[i].val - V[j].val;
496 if (deltaMin > delta)
504 if( DEBUG_LEVEL > 3 ) {
505 printf(
"deltaMin=%f\n", deltaMin);
510 fprintf(stderr,
"emd: Unexpected error in isOptimal.\n");
512 throw(
"EarthMoversDistance: Unexpected error in isOptimal.");
519 return deltaMin >= -
EPSILON * _maxC;
531 template<
class T,
class F>
537 node2_t *Loop[2*MAX_SIG_SIZE1], *CurX, *LeaveX;
539 if( DEBUG_LEVEL > 3 ) {
540 printf(
"EnterX = (%d,%d)\n", _EnterX->i, _EnterX->j);
547 _EnterX->NextC = _RowsX[i];
548 _EnterX->NextR = _ColsX[j];
554 steps = findLoop(Loop);
558 for (k=1; k < steps; k+=2)
560 if (Loop[k]->val < xMin)
568 for (k=0; k < steps; k+=2)
570 Loop[k]->val += xMin;
571 Loop[k+1]->val -= xMin;
574 if( DEBUG_LEVEL >= 3 ) {
575 printf(
"LeaveX = (%d,%d)\n", LeaveX->i, LeaveX->j);
582 if (_RowsX[i] == LeaveX)
583 _RowsX[i] = LeaveX->NextC;
585 for (CurX=_RowsX[i]; CurX != NULL; CurX = CurX->NextC)
586 if (CurX->NextC == LeaveX)
588 CurX->NextC = CurX->NextC->NextC;
591 if (_ColsX[j] == LeaveX)
592 _ColsX[j] = LeaveX->NextR;
594 for (CurX=_ColsX[j]; CurX != NULL; CurX = CurX->NextR)
595 if (CurX->NextR == LeaveX)
597 CurX->NextR = CurX->NextR->NextR;
610 template<
class T,
class F>
614 node2_t **CurX, *NewX;
615 char IsUsed[2*MAX_SIG_SIZE1];
617 for (i=0; i < _n1+_n2; i++)
621 NewX = *CurX = _EnterX;
622 IsUsed[_EnterX-_X] = 1;
630 NewX = _RowsX[NewX->i];
631 while (NewX != NULL && IsUsed[NewX-_X])
637 NewX = _ColsX[NewX->j];
638 while (NewX != NULL && IsUsed[NewX-_X] && NewX != _EnterX)
650 if( DEBUG_LEVEL > 3 ) {
651 printf(
"steps=%d, NewX=(%d,%d)\n", steps, NewX->i, NewX->j);
666 }
while (NewX != NULL && IsUsed[NewX-_X]);
670 IsUsed[*CurX-_X] = 0;
674 }
while (NewX == NULL && CurX >= Loop);
676 if( DEBUG_LEVEL > 3 ) {
677 printf(
"BACKTRACKING TO: steps=%d, NewX=(%d,%d)\n",
678 steps, NewX->i, NewX->j);
680 IsUsed[*CurX-_X] = 0;
684 }
while(CurX >= Loop);
688 fprintf(stderr,
"emd: Unexpected error in findLoop!\n");
691 if( DEBUG_LEVEL > 3 ) {
692 printf(
"FOUND LOOP:\n");
693 for (i=0; i < steps; i++)
694 printf(
"%d: (%d,%d)\n", i, Loop[i]->i, Loop[i]->j);
705 template<
class T,
class F>
708 int i, j, found, minI, minJ;
709 double deltaMin, oldVal, diff;
710 double Delta[MAX_SIG_SIZE1][MAX_SIG_SIZE1];
711 node1_t Ur[MAX_SIG_SIZE1], Vr[MAX_SIG_SIZE1];
712 node1_t uHead, *CurU, *PrevU;
713 node1_t vHead, *CurV, *PrevV;
714 node1_t *PrevUMinI, *PrevVMinJ, *Remember;
717 uHead.Next = CurU = Ur;
718 for (i=0; i < _n1; i++)
725 (--CurU)->Next = NULL;
727 vHead.Next = CurV = Vr;
728 for (j=0; j < _n2; j++)
735 (--CurV)->Next = NULL;
738 for(i=0; i < _n1 ; i++)
739 for(j=0; j < _n2 ; j++)
750 for(i=0; i < _n1 ; i++)
751 for(j=0; j < _n2 ; j++)
752 Delta[i][j] = _C[i][j] - Ur[i].val - Vr[j].val;
757 if( DEBUG_LEVEL > 3 ) {
759 for(CurU = uHead.Next; CurU != NULL; CurU = CurU->Next)
760 printf(
"[%d]",CurU-Ur);
763 for(CurV = vHead.Next; CurV != NULL; CurV = CurV->Next)
764 printf(
"[%d]",CurV-Vr);
773 for (CurU=uHead.Next; CurU != NULL; CurU=CurU->Next)
778 for (CurV=vHead.Next; CurV != NULL; CurV=CurV->Next)
782 if (deltaMin > Delta[i][j])
784 deltaMin = Delta[i][j];
800 Remember = PrevUMinI->Next;
801 addBasicVariable(minI, minJ, S, D, PrevUMinI, PrevVMinJ, &uHead);
804 if (Remember == PrevUMinI->Next)
806 for (CurV=vHead.Next; CurV != NULL; CurV=CurV->Next)
810 if (CurV->val == _C[minI][j])
815 for (CurU=uHead.Next; CurU != NULL; CurU=CurU->Next)
819 if (CurV->val <= _C[i][j])
820 CurV->val = _C[i][j];
824 diff = oldVal - CurV->val;
825 if (fabs(diff) <
EPSILON * _maxC)
826 for (CurU=uHead.Next; CurU != NULL; CurU=CurU->Next)
827 Delta[CurU->i][j] += diff;
833 for (CurU=uHead.Next; CurU != NULL; CurU=CurU->Next)
837 if (CurU->val == _C[i][minJ])
842 for (CurV=vHead.Next; CurV != NULL; CurV=CurV->Next)
846 if(CurU->val <= _C[i][j])
847 CurU->val = _C[i][j];
851 diff = oldVal - CurU->val;
852 if (fabs(diff) <
EPSILON * _maxC)
853 for (CurV=vHead.Next; CurV != NULL; CurV=CurV->Next)
854 Delta[i][CurV->i] += diff;
858 }
while (uHead.Next != NULL || vHead.Next != NULL);
866 template<
class T,
class F>
868 node1_t *PrevUMinI, node1_t *PrevVMinJ,
873 if (fabs(S[minI]-D[minJ]) <=
EPSILON * _maxW)
879 else if (S[minI] < D[minJ])
893 _IsX[minI][minJ] = 1;
898 _EndX->NextC = _RowsX[minI];
899 _EndX->NextR = _ColsX[minJ];
900 _RowsX[minI] = _EndX;
901 _ColsX[minJ] = _EndX;
905 if (S[minI] == 0 && UHead->Next->Next != NULL)
906 PrevUMinI->Next = PrevUMinI->Next->Next;
908 PrevVMinJ->Next = PrevVMinJ->Next->Next;
916 template<
class T,
class F>
924 if( DEBUG_LEVEL > 2 ) {
925 printf(
"SIG1\tSIG2\tFLOW\tCOST\n");
927 for(P=_X; P < _EndX; P++)
928 if (P != _EnterX && _IsX[P->i][P->j])
930 if( DEBUG_LEVEL > 2 ) {
931 printf(
"%d\t%d\t%f\t%f\n", P->i, P->j, P->val, _C[P->i][P->j]);
933 totalCost += (double)P->val * _C[P->i][P->j];
936 printf(
"COST = %f\n", totalCost);
940 #endif // REDUKT_EARTHMOVERSDISTANCE_H
#define EPSILON
Definition: EarthMoversDistance.h:26
C++ Template Wrapper for original Ansi C code from Yossi Rubner.
Definition: EarthMoversDistance.h:69
virtual float distance(const F &feature1, const F &feature2) const =0
float value() const
Returns value of optimal solution.
Definition: EarthMoversDistance.h:77
EarthMoversDistance(const T &signature1, const T &signature2, const GroundDistance< F > &dist)
Compute EMD between two signatures according to given GroundDistance.
Definition: EarthMoversDistance.h:171
virtual ~GroundDistance()
Definition: EarthMoversDistance.h:16
Interface for distance functions to use with EarthMoversDistance.
Definition: EarthMoversDistance.h:14
#define EMD_INFINITY
Definition: EarthMoversDistance.h:25