o
    ~Nha                     @   sp   d dl mZmZmZ d dlZd dlm  mZ d dl	Z
d dlmZmZ d dlZdZdZdd ZG dd	 d	ZdS )
    )AnyDictTupleN)defaultdictCounteri>  i   c                 c   s.    t dt| |D ]}| |||  V  q	dS )z>Helper to yield successive chunks of given size from iterable.r   N)rangelen)iterablesizei r   /app/shazam.pychunked
   s   r   c                	   @   sh   e Zd Zeedddddddf	defdd	ZddefddZdd Z	dd Z
dd Zdd ZdddZdS )	GPUShazamN   g?cudai  g?min_match_ratioc
           
      C   sz   || _ || _|p|d | _|| _|| _tj| j|d| _i | _d| _	d| _
d| _d| _|| _|| _|| _|	| _t | _dS )u   
        GPU-based Shazam with optional query‐match threshold.
        Args:
          …  
          min_query_matches: only accept a match if vote count ≥ this threshold
           )devicer      i   N)sample_raten_fft
hop_length	fan_valuer   torchhann_windowwindowindexmin_hash_time_deltamax_hash_time_deltamin_hash_freq_deltamax_hash_freq_deltakernal_sizer   min_prune_sizemax_prune_ratioset_set)
selfr   r   r   r   r   r   r#   r$   r%   r   r   r   __init__   s    	zGPUShazam.__init__Fdelta_compressc              	      s   t j|t j| jd}t j|| j| j| j| jddd}| }| j	d }t
j|dd| j	| j	fd||fd}d	}|dd|k|dd|k@ }	|	   }	t|	\}
}tt||
}t|dkrlg S |jd
d d t| \}}g }tt|D ]^}td| jD ]U}|| t|k r|| || }}|||  |||  }}|| }|| }| j|  kr| jkrn q| j|  kr| jkrn q|d> |d> B |B }|||f qq|s|S tt |D ]\}} | | q fdd D S )z
        Compute fingerprints for a single audio waveform.
        
        Args:
            waveform (1D array-like): audio samples
            
        Returns:
            List of (hash, time_offset) tuples.
        )dtyper   FT)r   r   
win_lengthr   centerreturn_complex   r      )kernel_sizestridepadding
   c                 S   s   | d S )Nr   r   )xr   r   r   <lambda>T   s    z'GPUShazam.fingerprint.<locals>.<lambda>)key      c                    s   g | ]}| | fqS r   r   ).0hgroupingr   r   
<listcomp>l   s    z)GPUShazam.fingerprint.<locals>.<listcomp>)r   tensorfloat32r   stftr   r   r   absr#   F
max_pool2d	unsqueezesqueezecpunumpynpwherelistzipr   sortr   r   r   r    r!   r"   appendr   )r(   waveformr*   sigspecZmagr3   Z
max_pooledZamp_minZpeaksZfreqstimesZ	peak_listZtimes_sortedZfreqs_sortedhashesr   jt1f1t2f2dtdfr;   tr   r<   r   fingerprint+   sR   

(8
zGPUShazam.fingerprintc                 C   sP   | j |dd}t|dkr| j| |D ]\}}| j|g ||f qdS )z
        Add a song to the index.
        
        Args:
            song_id (Any): unique identifier for the song
            waveform (1D array-like): audio samples
        Tr*   r   N)r\   r   r'   addr   
setdefaultrN   )r(   song_idrO   Zfptr;   r[   r   r   r   add_songn   s   zGPUShazam.add_songc                 C   s   i }|D ]I\}}| j |g D ]=\}}t|tr4|D ]}|| }	|||	fd |||	f  d7  < qq|}|| }	|||	fd |||	f  d7  < qqt|}
|
|S )Nr   r0   )r   get
isinstancerK   r_   r   most_common)r(   batchtop_ncountsr;   Zt_queryr`   Z
t_ref_listt_refdeltacounterr   r   r   top_candidates|   s    

zGPUShazam.top_candidatesc                 C   sP   t | j| jk r
d S t | j| j }| jD ]}t | j| |kr%g | j|< qd S N)r   r'   r$   r%   r   )r(   Zmaximum_cnthashr   r   r   _prune   s   

zGPUShazam._prunec                 C   s   t |}||S rl   )r   rd   )r(   Zcounter_dictrf   rj   r   r   r   rd      s   
zGPUShazam.most_commonr0   c                 C   s   	 t | j|dd}t|}|dkrdg| S i }d}i d}}t||D ]b}	td|t|	   | |	|}
g d}}|
D ]1\}}||d}|rO|| nd}|| |t|	  }|| jkrjt|| dkrjd}|| ||< q?|r| 	|||t|	 f  S |t|	7 }q%| 	|||fS )	z
        Returns (best_song_id, best_time_delta, match_ratio).
        If no match or ratio < min_match_ratio, song_id/time_delta will be None.
        Fr]   r   )NNg        i  zProcessing batches in query, Tg{Gz?)
rK   r\   r   r   printrk   rb   r   rB   rd   )r(   rO   rf   rS   Ztotal_hashesrg   
batch_size
candidatestotalre   rk   resultsZshould_stoptpZ
vote_countZoriginal_countZoriginal_ratioZnew_match_ratior   r   r   query   s.   


zGPUShazam.query)F)r0   )__name__
__module____qualname__SAMPLE_RATEWINDOW_SIZEfloatr)   boolr\   ra   rk   rn   rd   ru   r   r   r   r   r      s    
Cr   )typingr   r   r   r   Ztorch.nn.functionalnn
functionalrC   rH   rI   collectionsr   r   heapqry   rz   r   r   r   r   r   r   <module>   s    