
    g:                         d dl Zd dlZd dlZd dlZd dlmZ d Zd ZddZ	ddZ
d Zd	 Z	 ddZd ZddZd Z G d d          Z G d de          Z G d d          Z G d d          ZdS )    N)
ThreadPoolc                       j         \  }}j         ||fk    sJ t           fdt          |          D                       }| j        z  S )z< computes the intersection measure of two result tables
    c              3   d   K   | ]*}t          j        |         |                   j        V  +d S N)npintersect1dsize).0iI1I2s     V/var/www/html/mpstechhub/venv/lib/python3.11/site-packages/faiss/contrib/evaluation.py	<genexpr>z+knn_intersection_measure.<locals>.<genexpr>   sN         	r!ube$$)         )shapesumranger	   )r   r   nqrankninters   ``   r   knn_intersection_measurer      sw     xHB8Dz!!!!     r    F BGr   c                    | j         dz
  }||k     }t          j        |           }t          |          D ]<}||         || |         | |dz                                                     z   ||dz   <   =|||         ||         fS )z select a set of results    )r	   r   
zeros_liker   r   )limsDIthreshr   masknew_limsr   s           r   filter_range_resultsr!      s    	QBv:D}T""H2YY J J"1+T!WtAE{-B(C(G(G(I(IIQQtWag%%r   overallc                 l   	
  fd
fd j         dz
  }j         dz
  |k    sJ t          j        |d          		
fd}t          d          }|                    |t          |                     t           dd          dd	         z
  dd         dd	         z
  	|
          S )zucompute the precision and recall of range search results. The
    function does not take the distances into account. c                 6    |          | dz                     S Nr    r   Ireflims_refs    r   ref_result_forz range_PR.<locals>.ref_result_for,       HQKQ/00r   c                 6    |          | dz                     S r%   r&   )r   Inewlims_news    r   new_result_forz range_PR.<locals>.new_result_for/   r+   r   r   int64dtypec                      |           } |           }t          j        ||          }t          |          | <   d S r   )r   r   len)qgt_idsnew_idsinterr/   r   r*   s       r   compute_PR_forz range_PR.<locals>.compute_PR_for7   sK      "" !.## vw//JJq			r      Nmode)r	   r   zerosr   mapr   counts_to_PR)r)   r(   r.   r-   r=   r   r9   poolr/   r   r*   s   ````    @@@r   range_PRrB   (   s   1 1 1 1 1 11 1 1 1 1 1 
	B=1""""Xb(((F       b>>DHH^U2YY'''x}$x}$	   r   c                 (   |dk    re|                                  |                                 |                                 }}} |dk    r||z  }nd}| dk    r|| z  }n|dk    rd}nd}||fS |dk    r| dk    }d| |<   || z  }||         dk                        t                    ||<   |dk    }t          j        ||         dk              sJ d||<   d||<   ||z  }|                                |                                fS t                      )z computes a  precision-recall for a ser of queries.
    ngt = nb of GT results per query
    nres = nb of found results per query
    ninter = nb of correct results per query (smaller than nres of course)
    r"   r         ?        averager   )r   astypefloatr   allmeanAssertionError)	ngtnresr   r=   	precisionrecallr   recalls
precisionss	            r   r@   r@   P   s8    yGGIItxxzz6::<<6T!88III77c\FFQYYFFF&  			 axD	3,dq0077 qyvfTla'(((((tT
d]
  ',,..00 r   c                 L   t          j        |          }t          j        |          }t          |           dz
  }t          |          D ]W}| |         | |dz            }}|||         }	|||         }
|
                                }|	|         |||<   |
|         |||<   X||fS )z& sort 2 arrays using the first as key r   )r   
empty_liker4   r   argsort)r   r   r   r   D2r   r   l0l1iidios               r   sort_range_res_2r[   ~   s    	q		B	q		B	TQB2YY  a$q1u+Br"uXr"uXJJLLqE2b5	qE2b5		r6Mr   c                     t          j        |          }t          |           dz
  }t          |          D ]@}| |         | |dz            }}|||         |||<   |||                                          A|S r%   )r   rS   r4   r   sort)r   r   r   r   r   rV   rW   s          r   sort_range_res_1r^      s    	q		B	TQB2YY  a$q1u+BbeH2b5	
2b5	Ir   ref,newc           	      |    d|v rt                     d|v rt                    \   fdfd j        dz
  }j        dz
  |k    sJ t                    }	t	          j        ||	dfd          fd	}
t          d
          }|                    |
t          |                     t	          j        |	          }t	          j        |	          }t          |	          D ]C}t          dd|df         dd|df         dd|df         |          \  }}|||<   |||<   D||fS )z compute precision-recall values for range search results
    for several thresholds on the "new" results.
    This is to plot PR curves
    refnewc                 6    |          | dz                     S r%   r&   r'   s    r   r*   z4range_PR_multiple_thresholds.<locals>.ref_result_for   r+   r   c                 R    |          | dz            }}||         ||         fS r%   r&   )r   rV   rW   Dnewr-   r.   s      r   r/   z4range_PR_multiple_thresholds.<locals>.new_result_for   s3    !hq1uoBBrE{DBK''r   r      r0   r1   c                     	|           } |           \  }}t          |          | d d df<   |j        dk    rd S t          j        |
          }|| d d df<   |j        dk    rd S t          j        ||          }d||t          |          k    <   t          j        ||         |k              }t          j        dg|f          }||         | d d df<   d S )Nr   r   r;      )r4   r	   r   searchsortedcumsumhstack)r5   r6   res_idsres_disrM   rX   n_okcountsr/   r*   
thresholdss          r   r9   z4range_PR_multiple_thresholds.<locals>.compute_PR_for   s    "")>!,,f++q!!!Qw<1F ogz22q!!!Qw;!F _VW-- "2Vyw.// y1#t%%t*q!!!Qwr   r:   Nr   rh   r<   )
r^   r[   r	   r4   r   r>   r   r?   r   r@   )r)   r(   r.   re   r-   rp   r=   do_sortr   ntr9   rA   rQ   rP   tprro   r/   r*   s   ``````           @@@r   range_PR_multiple_thresholdsrv      s    $// %hd;;
d1 1 1 1 1 1( ( ( ( ( ( ( 
	B=1""""	ZBXr2qk111F% % % % % % % %4 b>>DHH^U2YY''' "JhrllG2YY  qqq!Qw1a&Aq/
 
 
1 
1

wr   c                 X   t          j        | |g          }|                                 t          |          }t          j        |          }|dd         |dd         z
  |dd<   |||k             }t          j        || d          dz
  }t          j        ||d          dz
  }||fS )zt for two tables, cluster them by merging values closer than thr.
    Returns the cluster ids for each table element r   Nr;   right)side)r   rk   r]   r4   onesri   )	tab1tab2thrtabndiffsunique_valsidx1idx2s	            r   _cluster_tables_with_tolerancer      s     )T4L
!
!CHHJJJCAGAJJEABB#crc("E!""Ieck"K?;7;;;a?D?;7;;;a?D:r   h㈵>c           
      F   t           j                            | ||           t          j                    }t          t          |                    D ]}t          j        ||         ||         k              r'|| |                                         z  }t          | |         ||         |          \  }}	t          j
        |          D ]U}
|
|d         k    r||
k    }|                    t          |||f                   t          |||f                              VdS )zS test that knn search results are identical, with possible ties.
    Raise if not. )rtolr;   N)r   testingassert_allcloseunittestTestCaser   r4   rI   maxr   uniqueassertEqualset)Drefr(   re   r-   r   testcaser   ru   DrefCDnewCdisr   s               r   check_ref_knn_with_drawsr      s&    JtT555 ""H3t99 I I6$q'T!W$%% 	 47;;== 5d1gtAwJJu9U## 	I 	ICeBiC<D  T!T']!3!3Sag5G5GHHHH		II Ir   c                    t           j                            | |           t          |           dz
  }t	          |          D ]}| |         | |dz            }	}|||	         }
|||	         }|||	         }|||	         }t          j        |
|k              rnAd } ||
|          \  }
} |||          \  }}t           j                            |
|           t           j                            ||d           dS )zM compare range search results wrt. a reference result,
    throw if it fails r   c                 J    |                                  }| |         ||         fS r   )rT   )r   r   rZ   s      r   sort_by_idsz,check_ref_range_results.<locals>.sort_by_ids  s!    IIKKtQqTz!r      )decimalN)r   r   assert_array_equalr4   r   rI   assert_array_almost_equal)Lrefr   r(   Lnewre   r-   r   r   rV   rW   Ii_refIi_newDi_refDi_newr   s                  r   check_ref_range_resultsr   	  s&    J!!$---	TQB2YY H Ha$q1u+Bbebebebe6&F"## 		:" " "  +{66::VV*{66::VVJ))&&999

,,VVQ,GGGG!H Hr   c                   <    e Zd ZdZd Zd Zd Zd Zd Zd Z	d Z
d	S )
OperatingPointszw
    Manages a set of search parameters with associated performance and time.
    Keeps the Pareto optimal points.
    c                 "    g | _         g | _        d S r   )operating_pointssuboptimal_pointsselfs    r   __init__zOperatingPoints.__init__,  s    !
 "$r   c                     t           )z1 return -1 if k1 > k2, 1 if k2 > k1, 0 otherwise NotImplementedr   k1k2s      r   compare_keyszOperatingPoints.compare_keys3      r   c                     t           )zC parameters to say we do noting, takes 0 time and has 0 performancer   r   s    r   do_nothing_keyzOperatingPoints.do_nothing_key7  r   r   c                 @    | j         D ]\  }}}||k    r	||k    r dS dS )NFT)r   )r   perf_newt_new_perfrs   s         r   is_pareto_optimalz!OperatingPoints.is_pareto_optimal;  s:    / 	 	JAtQxAJJuutr   c                     d}d}| j         | j        z   D ]8\  }}}|                     ||          }|dk    r||k    r|}|dk     r||k     r|}9||fS )z, predicts the bound on time and performance rE   rD   r   )r   r   r   )r   keymin_timemax_perfkey2r   rs   cmps           r   predict_boundszOperatingPoints.predict_boundsA  s{    !2T5KK 	$ 	$MD$##C..CQwwx<< HQww(??#H!!r   c                 ^    |                      |          \  }}|                     ||          S r   )r   r   )r   r   r   r   s       r   should_run_experimentz%OperatingPoints.should_run_experimentO  s0    #223778%%h999r   c                    |                      ||          rd}|t          | j                  k     rm| j        |         \  }}}||k    r9||k     r3| j                            | j                            |                     n|dz  }|t          | j                  k     m| j                            |||f           dS | j                            |||f           dS )Nr   r   TF)r   r4   r   r   appendpop)r   r   r   rs   r   op_Lsperf2t2s           r   add_operating_pointz#OperatingPoints.add_operating_pointS  s    !!$** 	Ac$/0000#'#8#; ub5==QVV*11-11!446 6 6 6 FA c$/0000 !((#tQ8884"))3a.9995r   N)__name__
__module____qualname____doc__r   r   r   r   r   r   r   r&   r   r   r   r   &  s         
$ $ $      " " ": : :    r   r   c                   V    e Zd ZdZd Zd Zd Zd Zd Ze	j
        fdZd Zd	 Zd
 ZdS )OperatingPointsWithRangesz
    Set of parameters that are each picked from a discrete range of values.
    An increase of each parameter is assumed to make the operation slower
    and more accurate.
    A key = int array of indices in the ordered set of parameters.
    c                 H    t                               |            g | _        d S r   )r   r   rangesr   s    r   r   z"OperatingPointsWithRanges.__init__m  s!      &&&r   c                 >    | j                             ||f           d S r   )r   r   )r   namevaluess      r   	add_rangez#OperatingPointsWithRanges.add_ranger  s"    D&>*****r   c                 n    t          j        ||k              rdS t          j        ||k              rdS dS )Nr   r;   r   )r   rI   r   s      r   r   z&OperatingPointsWithRanges.compare_keysu  s=    6"( 	16"( 	2qr   c                 \    t          j        t          | j                  t                    S )Nr1   )r   r>   r4   r   intr   s    r   r   z(OperatingPointsWithRanges.do_nothing_key|  s!    xDK((4444r   c                 b    t          t          j        d | j        D                                 S )Nc                 2    g | ]\  }}t          |          S r&   )r4   )r
   r   r   s      r   
<listcomp>z=OperatingPointsWithRanges.num_experiments.<locals>.<listcomp>  s"    HHHLD&CKKHHHr   )r   r   prodr   r   s    r   num_experimentsz)OperatingPointsWithRanges.num_experiments  s+    27HHDKHHHIIJJJr   c                 6   |dk    s|dk    sJ |                                  }t          j                            d          }|dk    s||k     r|                    |dz
            }n|                    |dz
  |dz
  d          }d|dz
  gd |D             z   }|S )z} sample a set of experiments of max size n_autotune
        (run all experiments in random order if n_autotune is 0)
        r   rh   {   F)r	   replacer   c                 2    g | ]}t          |          d z   S )r   )r   )r
   cnos     r   r   z@OperatingPointsWithRanges.sample_experiments.<locals>.<listcomp>  s"    'L'L'LC1'L'L'Lr   )r   r   randomRandomStatepermutationchoice)r   
n_autotunerstotexexperimentss        r   sample_experimentsz,OperatingPointsWithRanges.sample_experiments  s     Q*////$$&&Y""3''??ej00..33KK))	
Q $ ? ?K %!)n'L'L'L'L'LLr   c                     t          j        t          | j                  t                    }t          | j                  D ]/\  }\  }}|t          |          z  ||<   |t          |          z  }0|dk    sJ |S )z/Convert a sequential experiment number to a keyr1   r   )r   r>   r4   r   r   	enumerate)r   r   kr   r   r   s         r   
cno_to_keyz$OperatingPointsWithRanges.cno_to_key  sz    HS%%S111!*4;!7!7 	  	 A~fV$AaDCKKCCaxxxxr   c                 D    fdt          | j                  D             S )z3Convert a key to a dictionary with parameter valuesc                 :    i | ]\  }\  }}|||                  S r&   r&   )r
   r   r   r   r   s       r   
<dictcomp>z<OperatingPointsWithRanges.get_parameters.<locals>.<dictcomp>  s;     
 
 
!>D& &1,
 
 
r   )r   r   )r   r   s    `r   get_parametersz(OperatingPointsWithRanges.get_parameters  s8    
 
 
 
%.t{%;%;
 
 
 	
r   c                     | j         D ]#\  }}||k    rfd|D             }||dd<    dS $t          d| d          )z% remove too large values from a rangec                      g | ]
}|k     |S r&   r&   )r
   vmax_vals     r   r   z<OperatingPointsWithRanges.restrict_range.<locals>.<listcomp>  s    999aQ[[[[[r   Nz
parameter z
 not found)r   RuntimeError)r   r   r   name2r   val2s     `   r   restrict_rangez(OperatingPointsWithRanges.restrict_range  so    ![ 	 	ME6u}}99996999 qqq	  8888999r   N)r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r&   r   r   r   r   e  s           
+ + +  5 5 5K K K 13	       
 
 
: : : : :r   r   c                       e Zd Zd Zd ZdS )	TimerIterc                     g | _         |j        | _        || _        |j        dk    rt	          j        |j                   d S d S )Nr   )tsrunstimerrr   faissomp_set_num_threads)r   r  s     r   r   zTimerIter.__init__  sD    J	
8q==%eh///// =r   c                 J   | j         }| xj        dz  c_        | j                            t	          j                               t          | j                  dk    r| j        d         | j        d         z
  nd}| j        dk    s||j        k    r|j        dk    rt          j	        |j
                   t          j        | j                  }|dd          |d d         z
  }t          |          |j        k    r||j        d          |_        n|d d          |_        t          d S )Nr   rh   r;   r   )r  r  r  r   timer4   max_secsrr   r  r  remember_ntr   arraywarmuptimesStopIteration)r   r  
total_timer  r  s        r   __next__zTimerIter.__next__  s   
		Q		ty{{###14TW1B1BTWR[471:--
9??j5>99x1}})%*;<<<$'""BqrrFRW$E5zzUZ''#ELNN3 $AAAh :9r   N)r   r   r   r   r  r&   r   r   r   r     s2        0 0 0         r   r   c                   D    e Zd ZdZdddej        fdZd Zd Zd Z	d	 Z
d
S )RepeatTimeru!  
    This is yet another timer object. It is adapted to Faiss by
    taking a number of openmp threads to set on input. It should be called
    in an explicit loop as:

    timer = RepeatTimer(warmup=1, nt=1, runs=6)

    for _ in timer:
        # perform operation

    print(f"time={timer.get_ms():.1f} ± {timer.get_ms_std():.1f} ms")

    the same timer can be re-used. In that case it is reset each time it
    enters a loop. It focuses on ms-scale times because for second scale
    it's usually less relevant to repeat the operation.
    r   r;   r   c                 ~    ||k     sJ || _         || _        || _        || _        t	          j                    | _        d S r   )r  rr   r  r  r  omp_get_max_threadsr	  )r   r  rr   r  r  s        r   r   zRepeatTimer.__init__  sB    }}}}	  466r   c                      t          |           S r   )r   r   s    r   __iter__zRepeatTimer.__iter__  s    r   c                 :    t          j        | j                  dz  S )N  )r   rJ   r  r   s    r   mszRepeatTimer.ms  s    wtz""T))r   c                 n    t          | j                  dk    rt          j        | j                  dz  ndS )Nr   r  rE   )r4   r  r   stdr   s    r   ms_stdzRepeatTimer.ms_std  s0    ,/
OOa,?,?rvdj!!D((SHr   c                 *    t          | j                  S )zJ effective number of runs (may be lower than runs - warmup due to timeout))r4   r  r   s    r   nrunszRepeatTimer.nruns  s    4:r   N)r   r   r   r   r   infr   r  r  r  r  r&   r   r   r  r    s~            BQ 7 7 7 7  * * *I I I    r   r  )r"   )r"   r_   )r   )numpyr   r   r  r  multiprocessing.poolr   r   r!   rB   r@   r[   r^   rv   r   r   r   r   r   r   r  r&   r   r   <module>r!     s          + + + + + +
	 	 	& & &% % % %P, , , ,\     %.	G G G G\  I I I I,H H H:< < < < < < < <~D: D: D: D: D: D: D: D:T               2$ $ $ $ $ $ $ $ $ $r   