
      iQ                         d dl Z d dlmZmZ d dl mZ ddlmZ ddlmZ dde j                  de j                  d	e	fd
Z
	 dde j                  de j                  d	e	fdZ G d de      Zy)    N)OptionalUnion)Tensor   )BaseWaveformTransform)
ObjectDicttensorrrolloverc                 .   | j                   \  }}}t        j                  || j                        }|ddddf   }||z
  j	                  |||g      }t        j
                  | d||z        }|r|S |dz   j                  d      }	d|	|	|kD  <   d||	dk(  <   |S )z Shift or roll a batch of tensors)deviceNr      r   )shapetorcharanger   expandgatherclamp)
r	   r
   r   bctxidxsret
cut_pointss
             ~/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torch_audiomentations/augmentations/shift.py	shift_gpur   	   s    llGAq! 	Qv}}-A 	
!T4-AE>>1a)$D
,,vq$(
+C
 (!!!$J!"JzA~C
aJ    selected_samplesshift_samplesc                     | j                  d      }t        |      D ]V  }||   j                         }t        j                  | |   |d      | |<   |r7|dkD  rd| |dd|f<   G|dk  sMd| |d|df<   X | S )zMShift or roll a batch of tensors with the help of a for loop and torch.roll()r   )shiftsdimsg        .N)sizerangeitemr   roll)r   r    r   selected_batch_sizeinum_samples_to_shifts         r   	shift_cpur,      s     +//2&',Q/446#jjQ(<2
 #a'BE C)>*>)>!>?%)BE C)=)>!>? ( r   c                       e Zd ZdZh dZdZdZdZdZ	 	 	 	 	 	 	 	 	 	 dde	e
ef   de	e
ef   deded	ed
e
dee   dee   dee   dee   f fdZ	 	 	 	 ddedee   dee   dee   fdZ	 	 	 	 ddedee   dee   dee   def
dZdefdZ xZS )ShiftzI
    Shift the audio forwards or backwards, with or without rollover
    >   	per_batchper_channelper_exampleTF	min_shift	max_shift
shift_unitr   modepp_modesample_ratetarget_rateoutput_typec                     t         |   |||||	|
       || _        || _        || _        || _        | j                  | j                  kD  rt        d      | j                  dvrt        d      y)aW  

        :param min_shift: minimum amount of shifting in time. See also shift_unit.
        :param max_shift: maximum amount of shifting in time. See also shift_unit.
        :param shift_unit: Defines the unit of the value of min_shift and max_shift.
            "fraction": Fraction of the total sound length
            "samples": Number of audio samples
            "seconds": Number of seconds
        :param rollover: When set to True, samples that roll beyond the first or last position
            are re-introduced at the last or first. When set to False, samples that roll beyond
            the first or last position are discarded. In other words, rollover=False results in
            an empty space (with zeroes).
        :param mode:
        :param p:
        :param p_mode:
        :param sample_rate:
        :param target_rate:
        )r5   r6   r7   r8   r9   r:   z,min_shift must not be greater than max_shift)fractionsamplessecondsz5shift_unit must be "samples", "fraction" or "seconds"N)super__init__r2   r3   r4   r   
ValueError)selfr2   r3   r4   r   r5   r6   r7   r8   r9   r:   	__class__s              r   r@   zShift.__init__@   s    > 	### 	 	
 #"$ >>DNN*KLL??"DDTUU Er   r=   targetsc                    | j                   dk(  r| j                  }| j                  }n| j                   dk(  r]t        t	        | j                  |j
                  d   z              }t        t	        | j                  |j
                  d   z              }n]| j                   dk(  rCt        t	        | j                  |z              }t        t	        | j                  |z              }nt        d      t        j                  t        j                        j                  |cxk  r1t        j                  t        j                        j                  k  sJ  J t        j                  t        j                        j                  |cxk  r1t        j                  t        j                        j                  k  sJ  J |j                  d      }||k(  r@t        j                  |f|t        j                  |j                        | j                  d<   y t        j                   ||d	z   |ft        j                  |j                  
      | j                  d<   y )Nr=   r<   r"   r>   zInvalid shift_unitr   )r%   
fill_valuedtyper   r+   r   )lowhighr%   rG   r   )r4   r2   r3   introundr   rA   r   iinfoint32minmaxr%   fullr   transform_parametersrandint)rB   r=   r8   rD   r9   min_shift_in_samplesmax_shift_in_samplesr)   s           r   randomize_parameterszShift.randomize_parametersp   s    ??i'#'>> #'>> __
*#&uT^^gmmB>O-O'P#Q #&uT^^gmmB>O-O'P#Q __	)#&uT^^k-I'J#K #&uT^^k-I'J#K  122 KK$((#,{{5;;'++,	
,	
,
 KK$((#,{{5;;'++,	
,	
, &ll1o#77@E

)+/kk~~	AD%%&<= AF()A-)+kk~~AD%%&<=r   returnc                 `   | j                   d   }|j                  j                  dk(  rt        nt        } |||| j
                        }||dk(  r|}nNt        t        ||z  |z              }	 ||j                  dd      |	| j
                        j                  dd      }t        ||||      S )Nr+   cudar   r"   )r=   r8   rD   r9   )
rQ   r   typer   r,   r   rJ   rK   	transposer   )
rB   r=   r8   rD   r9   r+   shiftshifted_samplesshifted_targetsnum_frames_to_shifts
             r   apply_transformzShift.apply_transform   s      $889OP %^^00F:		)=t}}M?kQ.%O #&k$88;FG# $!!"b)+>iB  ####	
 	
r   c                      | j                   dk(  S )Nr>   )r4   )rB   s    r   is_sample_rate_requiredzShift.is_sample_rate_required   s    )++r   )
g            ?r<   Tr1   rc   NNNN)NNNN)__name__
__module____qualname____doc__supported_modessupports_multichannelrequires_sample_ratesupports_targetrequires_targetr   floatrJ   strboolr   r@   r   rU   r   r`   rb   __classcell__)rC   s   @r   r.   r.   3   s~    BO OO (,'*$! $%)%)%).V$.V $.V 	.V
 .V .V .V .V c].V c].V c].Vd %)$(%)00 c]0 &!	0
 c]0h %)$(%)

 c]
 &!	

 c]
 

>, ,r   r.   )F)r   typingr   r   r   core.transforms_interfacer   utils.object_dictr   ro   r   r,   r.    r   r   <module>ru      sq     "  = *ell u|| t , SXll38<<KO*P,! P,r   