
    hJJjNo                     x   d dl Z d dlmZmZmZ d dlmZmZmZm	Z	m
Z
mZmZ d dlmZ d dlZd dlmZ ddlmZ ddlmZmZ defd	Z	 d,d
ddej6                  dedeej6                  ee   f   fdZ ed       G d d             Z ed       G d d             Z G d d      Z  G d d      Z! G d de!      Z" G d d      Z#ejH                  d        Z% G d de#      Z& G d  d!      Z' G d" d#e'      Z( G d$ d%e'      Z) G d& d'e'      Z* G d( d)      Z+ e       fd
ddej6                  d*edeeee   f   fd+Z,y)-    N)	dataclassfieldreplace)DictIterableListOptionalSequenceTupleUnion)tree_map   )CHUNK_LENGTH)	Tokenizerget_tokenizerreturnc                 x    | j                  d      }t        |      t        t        j                  |            z  S )Nzutf-8)encodelenzlibcompress)text
text_bytess     ^/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_whisper/decoding.pycompression_ratior      s-    W%Jz?Sz!:;;;    modelWhispermel	tokenizerc                 D   |!t        | j                  | j                        }|j                  |j                  |j
                  vrt        d      |j                  dk(  }|r|d   }|j                  dd | j                  j                  | j                  j                  fk7  r| j                  |      }|j                  d   }t        j                  |j                  gg|z        }| j!                  ||      dddf   }t        j"                  |j                  d   t        j$                   t        j&                        }d	|t)        |j*                        <   ||z  }t        j,                  |d
      }t        j.                  |d
      }	t1        j                  |	      }	t3        |      D 
cg c]I  }
t5        |j*                  |j6                        D ci c]  \  }}||	|
|f   j9                          c}}K }}}
}|r
|d   }|d   }||fS c c}}w c c}}}
w )aq  
    Detect the spoken language in the audio, and return them as list of strings, along with the ids
    of the most probable language tokens and the probability distribution over all language tokens.
    This is performed outside the main decode loop in order to not interfere with kv-caching.

    Returns
    -------
    language_tokens : mx.array, shape = (n_audio,)
        ids of the most probable language tokens, which appears after the startoftranscript token.
    language_probs : List[Dict[str, float]], length = n_audio
        list of dictionaries containing the probability distribution over all languages.
    N)num_languageszCThis model doesn't have language tokens so it can't perform lang id   r   )dtype        axis)r   is_multilingualr"   languagelanguage_tokensot_sequence
ValueErrorndimshapedimsn_audio_ctxn_audio_stateencodermxarraysotlogitsfullinffloat32listall_language_tokensargmaxsoftmaxnprangezipall_language_codesitem)r   r   r    singlen_audioxr8   masklanguage_tokenslanguage_token_probsijclanguage_probss                 r   detect_languagerO      s    !!!1D1D
	 	"##9+A+AAQ
 	
 XX]F$i yy~%**00%**2J2JKKmmC  iilG
9==/"W,-A\\!S!!Q$'F 776<<#bffWBJJ?D03Di++	,-
dNFiiR0O::f2688$89 w
  A I999;W;WX	
X1 #AqD)..00X	
     )!,'*N**	
s   3(H H;HHT)frozenc                   L   e Zd ZU dZeed<   dZee   ed<   dZe	ed<   dZ
ee   ed<   dZee   ed<   dZee   ed	<   dZee	   ed
<   dZee	   ed<   dZeeeee   f      ed<   dZeeeee   f      ed<   dZeeeee   f      ed<   dZeed<   dZeed<   dZee	   ed<   dZeed<   y)DecodingOptions
transcribetaskNr+   r'   temperature
sample_lenbest_of	beam_sizepatiencelength_penaltypromptprefixz-1suppress_tokensTsuppress_blankFwithout_timestampsg      ?max_initial_timestampfp16)__name__
__module____qualname__rT   str__annotations__r+   r	   rU   floatrV   intrW   rX   rY   rZ   r[   r   r   r\   r]   r   r^   boolr_   r`   ra    r   r   rR   rR   R   s     D# #Hhsm" K $J$!GXc]!#Ix}# $Hhuo$ '+NHUO* /3FHU3S	>*+2.2FHU3S	>*+2 <@OXeC#$678?ND  %$-08E?0 D$r   rR   c                      e Zd ZU ej                  ed<   eed<   dZee	ee
f      ed<    ee      Zee   ed<   dZeed<   ej$                  Ze
ed	<   ej$                  Ze
ed
<   ej$                  Ze
ed<   ej$                  Ze
ed<   y)DecodingResultaudio_featuresr+   NrN   )default_factorytokens r   avg_logprobno_speech_probrU   r   )rb   rc   rd   r5   r6   rf   re   rN   r	   r   rg   r   r<   ro   r   rh   r   r@   nanrq   rr   rU   r   rj   r   r   rl   rl   w   s~    HHM15NHT#u*-.5d3FDI3D#NKFFNE"K!vvu%r   rl   c                   p    e Zd Zd	dZdej
                  dej
                  dej
                  fdZd Zd Zy)
	Inferencec                      || _         d | _        y N)r   kv_cache)selfr   s     r   __init__zInference.__init__   s     %
r   ro   rm   r   c                     | j                   j                  ||| j                        \  }| _        }|j                  t        j
                        S )zAPerform a forward pass on the decoder and return per-token logitsrx   )r   decoderrx   astyper5   r;   )ry   ro   rm   r8   _s        r   r8   zInference.logits   sD    #'::#5#5NT]] $6 $
 q }}RZZ((r   c                     t        t        t                          k7  rt        fd| j                        | _        yy)z9Update the key-value cache according to the updated beamsc                     |    S rw   rj   )rG   source_indicess    r   <lambda>z.Inference.rearrange_kv_cache.<locals>.<lambda>   s
    q/@r   N)r<   rA   r   r   rx   )ry   r   s    `r   rearrange_kv_cachezInference.rearrange_kv_cache   s4     T%N(;"<==$%@$--PDM >r   c                     d | _         y rw   r|   ry   s    r   resetzInference.reset   s	    r   N)r   r   )	rb   rc   rd   rz   r5   r6   r8   r   r   rj   r   r   ru   ru      s8    )RXX )rxx )BHH )Qr   ru   c                   R    e Zd Zdeeej
                        deee      dee   fdZy)SequenceRankerro   sum_logprobsr   c                     t         )z
        Given a list of groups of samples and their cumulative log probabilities,
        return the indices of the samples in each group to select as the final result
        NotImplementedErrorry   ro   r   s      r   rankzSequenceRanker.rank   s
     "!r   N)	rb   rc   rd   r   r5   r6   rg   rh   r   rj   r   r   r   r      s9    "4>*":>tE{:K"	c"r   r   c                   P    e Zd ZdZdee   fdZdeeee         deee      fdZ	y)MaximumLikelihoodRankerz
    Select the sample with the highest log probabilities, penalized using either
    a simple length normalization or Google NMT paper's length penalty
    rZ   c                     || _         y rw   )rZ   )ry   rZ   s     r   rz   z MaximumLikelihoodRanker.__init__   s
    ,r   ro   r   c           
            fd}|D cg c]  }|D cg c]  }t        |       c} }}}t        ||      D cg c]!  \  }}t        j                   |||            # c}}S c c}w c c}}w c c}}w )Nc                     g }t        | |      D ]=  \  }}j                  |}nd|z   dz  j                  z  }|j                  ||z         ? |S )N      )rB   rZ   append)logprobslengthsresultlogproblengthpenaltyry   s         r   scoresz,MaximumLikelihoodRanker.rank.<locals>.scores   sa    F#&x#9&&.$G !"F
a/D4G4GGGg/0 $: Mr   )r   rB   r@   r>   )	ry   ro   r   r   str   pls	   `        r   r   zMaximumLikelihoodRanker.rank   sm    		 1771A&AqCFA&747g4NO4NDAq		&A,'4NOO '7Os   	A-A(A-&A3(A-N)
rb   rc   rd   __doc__r	   rg   rz   r   rh   r   rj   r   r   r   r      sC    
-x -P4T#Y0 PT%[@Q Pr   r   c                      e Zd Zd Zdej
                  dej
                  dej
                  deej
                  eej
                  f   fdZdej
                  dej
                  dee	e	ej
                        e
e
e      f   fdZy)	TokenDecoderc                      y)z=Initialize any stateful variables for decoding a new sequenceNrj   r   s    r   r   zTokenDecoder.reset   s    r   ro   r8   r   r   c                     t         )a  Specify how to select the next token, based on the current trace and logits

        Parameters
        ----------
        tokens : mx.array, shape = (n_batch, current_sequence_length)
            all tokens in the context so far, including the prefix and sot_sequence tokens

        logits : mx.array, shape = (n_batch, vocab_size)
            per-token logits of the probability distribution at the current step

        sum_logprobs : mx.array, shape = (n_batch)
            cumulative log probabilities for each sequence

        Returns
        -------
        tokens : mx.array, shape = (n_batch, current_sequence_length + 1)
            the tokens, appended with the selected next token

        completed : bool
            True if all sequences has reached the end of text

        sum_logprobs: mx.array, shape = (n_batch)
            updated cumulative log probabilities for each sequence

        r   )ry   ro   r8   r   s       r   updatezTokenDecoder.update   s
    8 "!r   c                     t         )a  Finalize search and return the final candidate sequences

        Parameters
        ----------
        tokens : mx.array, shape = (n_audio, n_group, current_sequence_length)
            all tokens in the context so far, including the prefix and sot_sequence

        sum_logprobs : mx.array, shape = (n_audio, n_group)
            cumulative log probabilities for each sequence

        Returns
        -------
        tokens : Sequence[Sequence[mx.array]], length = n_audio
            sequence of mx.arrays containing candidate token sequences, for each audio input

        sum_logprobs : List[List[float]], length = n_audio
            sequence of cumulative log probabilities corresponding to the above

        r   r   s      r   finalizezTokenDecoder.finalize   s
    , "!r   N)rb   rc   rd   r   r5   r6   r   ri   r   r
   r   rg   r   rj   r   r   r   r      s    L"hh"(*"@B"	rxxrxx'	("<"hh".0hh"	x*+T$u+->>	?"r   r   c                 F    t         j                  j                  | |z        S rw   )r5   randomcategorical)r8   temps     r   r   r      s    99  $//r   c                       e Zd ZdedefdZdej                  dej                  dej                  deej                  e	ej                  f   fdZ
dej                  dej                  fd	Zy
)GreedyDecoderrU   eotc                      || _         || _        y rw   )rU   r   )ry   rU   r   s      r   rz   zGreedyDecoder.__init__   s    &r   ro   r8   r   r   c                 &   | j                   dk(  r|j                  d      }nt        || j                         }|t        j                  |dd      z
  }|t        j
                  |j                  d         |f   }|||d d df   | j                  k7  z  z  }|d d df   | j                  k(  }|d|z
  z  | j                  |z  z   }t        j                  ||d d d f   gd      }t        j                  |d d df   | j                  k(        }|||fS )Nr   r%   r(   Tr)   keepdimsr   )
rU   r>   r   r5   	logsumexparanger0   r   concatenateall)	ry   ro   r8   r   next_tokensr   current_logprobseot_mask	completeds	            r   r   zGreedyDecoder.update  s
    q  --R-0K%fd.>.>?KBLLb4HH#BIIhnnQ.?$@+$MN(F1b5MTXX,EFF!R%=DHH,!Q\2TXX5HHQW)= >RHFF6!R%=DHH45	y,..r   c                 R    t        j                  |g d| j                        }||fS )N)r   r   r   )r   r   )constant_values)r5   padr   r   s      r   r   zGreedyDecoder.finalize  s$     8$((S|##r   N)rb   rc   rd   rg   rh   rz   r5   r6   r   ri   r   r   rj   r   r   r   r      su    E  /hh/(*/@B/	rxxrxx'	(/($rxx $rxx $r   r   c                   \    e Zd Zdej                  dej                  dej                  fdZy)LogitFilterr8   ro   r   c                     t         )a  Apply any filtering or masking to logits

        Parameters
        ----------
        logits : mx.array, shape = (n_batch, vocab_size)
            per-token logits of the probability distribution at the current step

        tokens : mx.array, shape = (n_batch, current_sequence_length)
            all tokens in the context so far, including the prefix and sot_sequence tokens

        r   ry   r8   ro   s      r   applyzLogitFilter.apply  s
     "!r   N)rb   rc   rd   r5   r6   r   rj   r   r   r   r     s(    "BHH "bhh "288 "r   r   c                   p    e Zd ZdededefdZdej                  dej                  dej                  fdZy	)
SuppressBlankr    sample_beginn_vocabc                     || _         t        j                  |t        j                        }t        j                   ||j                  d      |j                  gz   <   t        j                  |      | _	        y )N )
r   r@   zerosr;   r:   r   r   r5   r6   rH   )ry   r    r   r   rH   s        r   rz   zSuppressBlank.__init__/  sS    (xx,9;Yc"imm_45HHTN	r   r8   ro   r   c                 \    |j                   d   | j                  k(  r|| j                  z   S |S )Nr   )r0   r   rH   r   s      r   r   zSuppressBlank.apply5  s,    <<?d///DII%%r   N)	rb   rc   rd   r   rh   rz   r5   r6   r   rj   r   r   r   r   .  sB    #) #3 # #BHH bhh 288 r   r   c                   r    e Zd Zdee   defdZdej                  dej                  dej                  fdZy)	SuppressTokensr]   r   c                     t        j                  |t         j                        }t         j                   |t	        |      <   t        j                  |      | _        y rw   )r@   r   r;   r:   r<   r5   r6   rH   )ry   r]   r   rH   s       r   rz   zSuppressTokens.__init__<  s:    xx,')vvgT/"#HHTN	r   r8   ro   r   c                      || j                   z   S rw   )rH   r   s      r   r   zSuppressTokens.applyA  s    		!!r   N)	rb   rc   rd   r
   rh   rz   r5   r6   r   rj   r   r   r   r   ;  s?    # # #
"BHH "bhh "288 "r   r   c                   v    e Zd Zdededee   fdZdej                  dej                  dej                  fdZ	y	)
ApplyTimestampRulesr    r   max_initial_timestamp_indexc                 .    || _         || _        || _        y rw   )r    r   r   )ry   r    r   r   s       r   rz   zApplyTimestampRules.__init__F  s     #(+F(r   r8   ro   r   c                    t        j                  |j                  t         j                        }| j                  j
                  ,t         j                   |d d | j                  j
                  f<   |j                         }t        t        |            D ]K  }||   | j                  d  }t        |      dk\  xr |d   | j                  j                  k\  }t        |      dk  xs |d   | j                  j                  k\  }|r[|r-t         j                   ||| j                  j                  d f<   n,t         j                   ||d | j                  j                  f<   t        |      D 	cg c]!  \  }}	|	| j                  j                  kD  s |# }
}}	t        |
      dkD  s|
d   }|r|r|dz  }t         j                   ||| j                  j                  |f<   N t        |d         | j                  k(  rzt         j                   |d d d | j                  j                  f<   | j                  @| j                  j                  | j                  z   }t         j                   |d d |dz   d f<   t        j                   |      }|t        j"                  |dd      z
  }|d d | j                  j                  d f   j#                  dd      }|d d d | j                  j                  f   j%                  dd      }t        j&                  ||kD  t        j                   |d d d | j                  j                  f         |d d d | j                  j                  f<   ||z   S c c}	}w )Nr   r%   r#   r$   r   Tr   )r@   r   r0   r;   r    no_timestampsr:   tolistrA   r   r   timestamp_beginr   	enumerater   r5   r6   r   maxwhere)ry   r8   ro   rH   kseqlast_was_timestamppenultimate_was_timestamprK   v
timestampslast_timestamplast_allowedr   timestamp_logprobmax_text_token_logprobs                   r   r   zApplyTimestampRules.applyP  s(   xxbjj1>>''357VVGDDNN0001 s6{#A)D--/0CCAK#b'T^^-K-K"K  C1IB4>>+I+I I & ",ACDDNN::<<=57VVGD0dnn00001 (n,daDNN4R4R0Rn   :" ",B%)B"a'NLNFF7Q66GGH1 $4 vay>T...9;D4dnn44445 //;NN22T5U5UU  02vvgQq(**+ xx~BLLb4HH$Q(F(F(H%HISSd T 
 "*!-Mt~~/M/M-M*M!N!R!Rd "S "
 57HH 66VVG4dnn444455
Q0$..00001
 }Is   .!M1M1N)
rb   rc   rd   r   rh   r	   rz   r5   r6   r   rj   r   r   r   r   E  sQ    GG G &.c]	G;BHH ;bhh ;288 ;r   r   c                   :   e Zd ZU eed<   eed<   eed<   ee   ed<   ddde	fdZ
de	d	e	fd
Zd	ee   fdZd	ee   fdZdej"                  fdZdej"                  dej"                  fdZdej"                  dej"                  fdZdej"                  d	ee   fdZy)DecodingTask	inferencesequence_rankerr}   logit_filtersr   r   optionsc                 ,   || _         |j                  xs d}t        |j                  |j                  ||j
                        }|| _        | j                  |      | _        |j                  xs |j                  xs d| _        |j                  j                  | _        |j                  xs |j                  j                  dz  | _        |j                   | _        | j                  j"                  r|j$                  | _        | j'                         | _        t+        | j(                        | _        | j(                  j/                  |j0                        | _        t5        |      | _        t9        |j:                        | _        |j                  t?        d      tA        |jB                  |jD                        | _#        g | _$        | j                  jJ                  rN| jH                  jM                  tO        | j                  | j,                  |j                  jP                               | j                  jR                  rG| jH                  jM                  tU        | jW                         |j                  jP                               |j"                  s~tX        |j                  jZ                  z  }d }|j\                  r"t_        | j                  j\                  |z        }| jH                  jM                  ta        || j,                  |             y y )Nen)r"   r+   rT   r   r#   z*Beam search decoder is not yet implemented)1r   r+   r   r*   r"   rT   r    _verify_optionsr   rX   rW   n_groupr1   
n_text_ctxn_ctxrV   r-   r_   #sot_sequence_including_notimestamps_get_initial_tokensinitial_tokensr   r   indexr7   	sot_indexru   r   r   rZ   r   r   r   rU   r   r}   r   r^   r   r   r   r]   r   _get_suppress_tokensr   r2   r`   roundr   )ry   r   r   r+   r    	precisionr   s          r   rz   zDecodingTask.__init__  sl   
##+t!!!--	
	 %.(,(<(<W(E#--EEA**//
&11OUZZ5J5Ja5O(1(>(><<** ) M MD*.*B*B*D!$T%8%8!9"1177	F #5)  7w7M7MN (%&RSS()<)<immLDL  <<&&%%dnnd.?.?ASAST <<''%%t88:EJJ<N<NO ))$uzz'='==I*.',,.3LL66B/+ %%#t002M *r   r   c                 N   |j                   |j                  t        d      |j                  dk(  r|j                  t        d      |j                  |j                   t        d      |j
                  ,d|j
                  cxk  rdk  st        d       t        d      |S )Nz-beam_size and best_of can't be given togetherr   z4best_of with greedy sampling (T=0) is not compatiblez'patience requires beam_size to be givenr   z8length_penalty (alpha) should be a value between 0 and 1)rX   rW   r.   rU   rY   rZ   )ry   r   s     r   r   zDecodingTask._verify_options  s    (W__-HLMM!#* !WXX'G,=,=,EFGG!!-'',1,WXX -WXXr   c                 b   t        | j                        }| j                  j                  x}rqt	        |t
              r,| j                  j                  d|j                         z         n|}| j                  "| j                  dz  | j                  z
  }|| d  }||z   }| j                  j                  x}rot	        |t
              r,| j                  j                  d|j                         z         n|}| j                  j                  g|| j                  dz  dz
   d  z   |z   }t        |      S )Nr   r#   r   )r<   r-   r   r\   
isinstancere   r    r   striprV   r   r[   sot_prevtuple)ry   ro   r\   prefix_tokensmax_prefix_lenr[   prompt_tokenss          r   r   z DecodingTask._get_initial_tokens  s(   d''(\\(((6( fc* %%cFLLN&:; 
 *!%q4??!B -~o.> ?m+F\\(((6( fc* %%cFLLN&:;  (()$**/A"5 6 89:  V}r   c                 &   | j                   j                  }t        |t              r'|j	                  d      D cg c]  }t        |       }}d|v r;|D cg c]
  }|dk\  s	| }}|j                  | j                  j                         n*|t        |      dk(  rg }nt        |t              sJ d       |j                  | j                  j                  | j                  j                  | j                  j                  | j                  j                  | j                  j                  g       | j                  j                   %|j#                  | j                  j                          t%        t'        t)        |                  S c c}w c c}w )N,r%   r   zsuppress_tokens must be a list)r   r]   r   re   splitrh   extendr    non_speech_tokensr   r<   rS   	translater7   r   sot_lm	no_speechr   r  sortedset)ry   r]   r   s      r   r   z!DecodingTask._get_suppress_tokens  sF   ,,66os+/>/D/DS/IJ/I!s1v/IOJ *9D/QQ!Vq/OD""4>>#C#CD$O(<(A Oot4V6VV4))((""''%%	
 >>##/""4>>#;#;<VC0122/ K Es   F	
F!Fr   c                    | j                   j                  r|j                  t        j                        }|j
                  dd  | j                  j                  j                  | j                  j                  j                  fk(  r|}n| j                  j                  |      }|j                  | j                   j                  rt        j                  nt        j                  k7  rt        d|j                         |S )Nr$   z'audio_features has an incorrect dtype: )r   ra   r~   r5   float16r0   r   r1   r2   r3   r4   r&   r;   	TypeError)ry   r   rm   s      r   _get_audio_featuresz DecodingTask._get_audio_features  s    <<**RZZ(C99RS>JJOO''JJOO))
 

 !N!ZZ//4N$,,2C2CBJJT9.:N:N9OP  r   rm   ro   c                    | j                   j                  g|j                  d   z  }d }| j                   j                  | j                   j                  dk(  r| j                  j                  || j                        \  }}|D cg c]  }t        ||j                         }}| j                   j                  )t        j                  |      |d d | j                  dz   f<   ||fS c c}w )Nr   lang_id)keyr   )r   r+   r0   rT   r   rO   r    r   getr@   r6   r   )ry   rm   ro   	languages
lang_probslang_tokensprobss          r   _detect_languagezDecodingTask._detect_language-  s    \\**+n.B.B1.EE	
<<  (DLL,=,=,J&*jj&@&@'#K AKK
uU		2
IK||$$,020Eq$..1,,-*$$ Ls   C)c                     |j                   d   }t        j                  |      } fd} |||||      \  }}}} j                  j                  Ft        j
                  |d d  j                  f   d      }|d d  j                  j                  f   }	n$t        j                  |t        j                        }	t        j                  ||||	       t        d j                        D ]\  }
|d d dd f   }|j                   d    j                  kD  r n3 |||||      \  }}}}t        j                  |||       |r n|}|}|}^ |||	fS )Nr   c                     j                   j                  | |      }|d d df   }j                  D ]  }|j                  ||      } j                  j                  |||      \  }}}||||fS )Nr%   )r   r8   r   r   r}   r   )	inputsrm   ro   r   
pre_logitsr8   logit_filterr   ry   s	           r   _stepz&DecodingTask._main_loop.<locals>._step@  s    ..v~FJ  2&F !% 2 2%++FF; !3 /3ll.A.A/+FI| 9lJ>>r   r%   r(   r   )r0   r5   r   r    r  r?   r   r9   rs   
async_evalrA   rV   r   )ry   rm   ro   n_batchr   r!  r   r  probs_at_sotno_speech_probsrK   r  r   next_completednext_sum_logprobsr   s   `               r   
_main_loopzDecodingTask._main_loop<  sA   ,,q/xx(	?  7<NFL7
3	< >>##/::jDNN1B&C"ML*1dnn.F.F+FGO gggrvv6O
iGq$//*AArsF^F||B$**,@EA=K):A MM.+7HI F&I,L + |_44r   c                    | j                   j                          | j                  j                          | j                  }|j                  d   }| j                  |      }t        j                  | j                        }t        j                  ||t        | j                        f      }| j                  ||      \  }}| j                  j                  dk(  r/t        |||      D 	
cg c]  \  }}	}
t        ||	|
       c}
}	}S | j                   dkD  ru|d d d d d f   }t        j                  ||| j                   t        | j                        g      }|j#                  || j                   z  t        | j                        f      }| j%                  ||      \  }}}|d d | j                      }|d d | j                      }|j                  d   t        |      cxk(  r|k(  sJ  J |j#                  || j                   d      }|j#                  || j                         }| j                  j'                  ||      \  }}|d| j(                  d f   }t        j*                  |||       |j-                         }|j-                         }|j-                         }|D cg c].  }|D cg c]   }|d |j/                  |j0                         " c}0 }}}| j2                  j5                  ||      }t        ||      D cg c]
  \  }}||    }}}|D cg c]!  }|j7                  |      j9                         # }}t        ||      D cg c]
  \  }}||    }}}t        ||      D cg c]  \  }}|t        |      dz   z   }}}||||||f}t        t;        t=        t        |                  dk7  r%t?        dtA        t=        t        |                   t        | D 	cg c]9  \  }}	}}}}t        ||	||||| j                  jB                  tE        |            ; c}}}}}	}S c c}
}	}w c c}w c c}}w c c}}w c c}w c c}}w c c}}w c c}}}}}	}w )	Nr   r  )rm   r+   rN   r   r%   .zinconsistent result lengths: )rm   r+   ro   r   rq   rr   rU   r   )#r   r   r}   r    r0   r  r5   r6   r   broadcast_tor   r  r   rT   rB   rl   r   reshaper(  r   r   evalr   r   r   r   r   decoder   r  mapRuntimeErrorr<   rU   r   )ry   r   r    rF   rm   ro   r  rN   featuresr+   r  r   r%  r   r   selectedrK   textslpavg_logprobsfieldsr   rq   rr   s                           r   runzDecodingTask.runk  sD   #~~	yy|#'#;#;C#@88D$7$78'3t7J7J3K)LM %)$9$9.&$Q!	><<	)
 25"I~2	2-Hh #+hu2	  <<!AtQJ'F__$,,D4G4G0HIF ^^Wt||%;SATAT=U$VWF 15PV0W-o (4<<8)/T\\/:##A&#o*>I'IIIIIr:#++GT\\B  $||44V\JT..001 	o6#**,)002DJKFq:11-qwwy}}-.:FK '',,V\B47&4I"J4IDAq1Q44I"JAGHAI,,Q/557H8;Hl8S$T8Suq"RU8S$T+.v|+D%
+D%!RB#a&1*+D 	 %

 
 s3sF#$%*!>tCVDT?U>VWXX RUR
 
RMh+~ '!'- LL44"3D"9	R
 	
yJ ;K #KH$T%

s<   *P7#	Q,%P>QQ	&QQ2Q1>Q >QN)rb   rc   rd   ru   rf   r   r   r   r   rR   rz   r   r   rh   r   r   r5   r6   r  r@   r  r(  rl   r6  rj   r   r   r   r     s    ##$$;i ;/ ;z ? U3Z 83eCj 3:rxx (%rxx % %-5 -5288 -5^W
rxx W
D$8 W
r   r   r   c                     |j                   dk(  x}r|d   }|rt        |fi |}t        | |      j                  |      }|r|d   S |S )a7  
    Performs decoding of 30-second audio segment(s), provided as Mel spectrogram(s).

    Parameters
    ----------
    model: Whisper
        the Whisper model instance

    mel: mx.array, shape = (80, 3000) or (*, 80, 3000)
        An array containing the Mel spectrogram(s)

    options: DecodingOptions
        A dataclass that contains all necessary options for decoding 30-second segments

    Returns
    -------
    result: Union[DecodingResult, List[DecodingResult]]
        The result(s) of decoding contained in `DecodingResult` dataclass instance(s)
    r#   Nr   )r/   r   r   r6  )r   r   r   kwargsrE   r   s         r   r-  r-    sZ    2 Qv$i',V,%)--c2F6!9*F*r   rw   )-r   dataclassesr   r   r   typingr   r   r   r	   r
   r   r   mlx.corecorer5   numpyr@   	mlx.utilsr   audior   r    r   r   rg   r   r6   dictrO   rR   rl   ru   r   r   r   compiler   r   r   r   r   r   r   r-  rj   r   r   <module>rB     s    1 1 I I I     /<u < =A;+;+88;+09;+
288T$Z ;+| $! ! !H $	& 	& 	& ," "Pn P48" 8"v 0 0$L $>" " 
K 
"[ "F+ FRt
 t
t	  /0 + +	 +  +
 >4//0 +r   