+
    TV-jX2  ã                   óX  € ^ RI t ^ RIHt ^ RIHtHtHtHt ^ RIH	t
 RRR^^ RR. 3R R lltRR R llt]! ]
P                  ]
P                  P                  ]
P                  P                  R	7      R
 R l4       t]! ]
P                  ]
P                  P                  ]
P                  P                  R	7      RR R ll4       t]! ]
P                  ]
P                  P                  ]
P                  P                  R	7      R R l4       t]! ]
P                  ]
P                  P                  ]
P                  P                  R	7      R R l4       t]! ]
P                  ]
P                  P                  ]
P                  P                  R	7      R 4       tRR R lltRR R lltRR R lltR# )é    N)Úpartial)ÚCallableÚDictÚListÚOptionalç        c                óæ   € V ^8„  d   QhR\         R\         R\         R\        R\        R\         R\         R\        \        ,          R	\        \        P
                  .\        P
                  3,          /	# )
é   ÚtempÚtop_pÚmin_pÚmin_tokens_to_keepÚtop_kÚxtc_probabilityÚxtc_thresholdÚxtc_special_tokensÚreturn)ÚfloatÚintr   r   ÚmxÚarray)Úformats   "Úd/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/mlx_lm/sample_utils.pyÚ__annotate__r   
   sz   € ÷ ;ñ ;Ü
ð;äð;ô ð;ô ð	;ô
 ð;ô ð;ô ð;ô œS�	ð;ô Œr�x‰xˆjœ"Ÿ(™(Ð"Õ#ñ;ó    c                ó2  a aaaaaaaa	€ S ^ 8X  d   R # . o	S^ 8”  d   SR8  d   S	P                  V3R l4       SR8w  d   S	P                  VV3R l4       SR8”  d   S	P                  VVV3R l4       S^ 8”  d   S	P                  V3R l4       V	V 3R lpV# )	a  
Make a sampler function for use with ``generate_step``.

Args:
    temp (float): The temperature for sampling, if 0 the argmax is used.
      Default: ``0``.
    top_p (float, optional): Nulceus sampling, higher means model considers
      more less likely words.
    min_p (float, optional): The minimum value (scaled by the top token's
      probability) that a token probability must have to be considered.
    min_tokens_to_keep (int, optional): Minimum number of tokens that cannot
      be filtered by min_p sampling.
    top_k (int, optional): The top k tokens ranked by probability to constrain
      the sampling to.
    xtc_probability (float, optional): The probability of applying XTC
        sampling.
    xtc_threshold (float, optional): The threshold the probs need to reach
        for being sampled.
    xtc_special_tokens (list(int), optional): List of special tokens IDs to
        be excluded from XTC sampling.


Returns:
    Callable[mx.array, mx.array]:
        A sampler which takes log-probabilities and returns tokens.
c                 ó2   € \         P                  ! V RR7      # )é   ©Úaxiséÿÿÿÿ)r   Úargmax)Úxs   &r   Ú<lambda>Úmake_sampler.<locals>.<lambda>/   s   € œŸš 1¨2Õ.r   ç      ð?c                 ó   <€ \        V S4      # ©N)Úapply_top_p)r#   r   s   &€r   r$   r%   4   ó   ø€ ¬+°a¸Ô*?r   r   c                 ó   <€ \        V SS4      # r(   )Úapply_min_p)r#   r   r   s   &€€r   r$   r%   6   s   ø€ ¬+°a¸Ð@RÔ*Sr   c                 ó    <€ \        V SSS4      # r(   )Ú	apply_xtc)r#   r   r   r   s   &€€€r   r$   r%   9   s   ø€ ”i  ?°MÐCUÔVr   c                 ó   <€ \        V S4      # r(   )Úapply_top_k)r#   r   s   &€r   r$   r%   <   r*   r   c                 ó>   <€ S F  pV! V 4      p K  	  \        V S4      # r(   )Úcategorical_sampling)ÚlogprobsÚmethodÚsampling_methodsr   s   & €€r   ÚsamplerÚmake_sampler.<locals>.sampler?   s&   ø€ Û&ˆFÙ˜hÓ'ŠHñ 'ô $ H¨dÓ3Ð3r   )Úappend)
r   r   r   r   r   r   r   r   r6   r5   s
   ffffffff @r   Úmake_samplerr9   
   s‹   ÿø€ ðH ˆq„yÙ.Ð.ð ÐØˆq„y�U˜S”[Ø×ÑÔ ?Ô@Ø�„|Ø×ÑÕ SÔTØ˜ÔØ×ÑÞVô	
ð ˆq„yØ×ÑÔ ?Ô@ö4ð €Nr   c                ó(  € V ^8„  d   QhR\         \        \        \        3,          ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          R\         \        ,          /# )r
   Ú
logit_biasÚrepetition_penaltyÚrepetition_context_sizeÚpresence_penaltyÚpresence_context_sizeÚfrequency_penaltyÚfrequency_context_size)r   r   r   r   )r   s   "r   r   r   H   st   € ÷ 6ñ 6Üœœc¤5˜jÕ)Õ*ð6ä ¤�ð6ô &¤c�]ð6ô œu•oð	6ô
 $¤C�=ð6ô  ¤•ð6ô %¤S�Mñ6r   c                ó”  aa€ . pV '       ds   \         P                  ! \        V P                  4       4      4      o\         P                  ! \        V P	                  4       4      4      oVV3R lpVP                  V4       \        W3\        W43\        WV3.p	V	 F,  w  r«pVf   K  V^ 8w  g   K  VP                  V
! W¼4      4       K.  	  V# )aQ  
Make logits processors for use with ``generate_step``.

Args:
    repetition_penalty (float, optional): A (sign-aware) multiplicative
      penalty for repeating tokens.
    repetition_context_size (int, optional): The number of tokens to
      consider for repetition penalty. Default: ``20``.
    presence_penalty (float, optional): An additive penalty to reduce
      repeating tokens.
    presence_context_size (int, optional): The number of tokens to consider
      for the presence penalty. Default: ``20``.
    frequency_penalty (float, optional): An additive penalty to reduce
      repeating tokens. The tokens are penalized proportionally to their
      frequency.
    frequency_context_size (int, optional): The number of tokens to consider
      for the frequency penalty. Default: ``20``.
    logit_bias (dictionary, optional): Additive logit bias.

Returns:
    List[Callable[[mx.array, mx.array], mx.array]]:
        A list of logits processors. Each processor in the list is a
        callable which takes an array of tokens and an array of logits
        and returns the updated logits.
c                 óL   <€ VP                   R S3,          P                  S4      # )ºNNN)ÚatÚadd)Ú_ÚlogitsÚindicesÚvaluess   &&€€r   Úlogit_bias_processorÚ4make_logits_processors.<locals>.logit_bias_processoro   s!   ø€ Ø—9‘9˜Q ˜ZÕ(×,Ñ,¨VÓ4Ð4r   )	r   r   ÚlistÚkeysrJ   r8   Úmake_repetition_penaltyÚmake_presence_penaltyÚmake_frequency_penalty)r;   r<   r=   r>   r?   r@   rA   Úlogits_processorsrK   Úrepetition_penaltiesÚmake_penaltyÚpenaltyÚcontext_sizerI   rJ   s   &&&&&&&      @@r   Úmake_logits_processorsrW   H   sº   ù€ ðD ÐßÜ—(’(œ4 
§¡Ó 1Ó2Ó3ˆÜ—’œ$˜z×0Ñ0Ó2Ó3Ó4ˆö	5ð 	× Ñ Ð!5Ô6ô 
!Ð"4ÐNÜ	Ð 0ÐHÜ	Ð!2ÐKðÐó 0DÑ+ˆ˜|ØÔ 7¨a¦<Ø×$Ñ$¡\°'Ó%HÖIñ 0Dð Ðr   )ÚinputsÚoutputsc                ód   € V ^8„  d   QhR\         P                  R\        R\         P                  /# )r
   r3   r   r   )r   r   r   )r   s   "r   r   r   ‚   s.   € ÷ ñ Ü�h‰hðäðô ‡X�Xñr   c           	     óz  € V P                   R	,          p\        V\        4      '       d   ^ Tu;8  d   V8  g   M \        RV RV R24      h\        P
                  ! V ) V^,
          R	R7      RVR13,          p\        P                  ! W\        P                  ! \        R4      ) V P                  4      R	R7      pV# )
zœ
Sample from only the top K tokens ranked by probability.

Args:
    logprobs: A vector of log probabilities.
    top_k (int): Top k tokens to sample from.
z(`top_k` has to be an integer in the (0, z] interval, but is Ú.©Úkthr    .NÚinfr   r!   )
ÚshapeÚ
isinstancer   Ú
ValueErrorr   ÚargpartitionÚput_along_axisr   r   Údtype)r3   r   Ú
vocab_sizeÚmask_idxÚmasked_logprobss   &&   r   r0   r0   �   sª   € ð —‘ Õ#€JÜ�eœS×!Ò!¨!¨eÖ*@°jÖ*@ÜØ6°z°lð CØ�g˜Qð ó
ð 	
ô �Š ˜y¨e°a­i¸bÔAÀ#ÀuÁvÀ+ÕN€HÜ×'Ò'ØœBŸHšH¤e¨E£l ]°H·N±NÓCÈ"ô€Oð Ðr   c                óp   € V ^8„  d   QhR\         P                  R\        R\        R\         P                  /# )r
   r3   r   r   r   )r   r   r   r   )r   s   "r   r   r   ›   s8   € ÷ .?ñ .?Ü�h‰hð.?äð.?ô ð.?ô ‡X�Xñ	.?r   c                óÞ  € ^ Tu;8:  d   R8:  g   M \        RV 24      h\        V\        4      '       d   V^8  d   \        RV 24      h\        P                  ! V RRR7      pV\
        P                  ! V4      ,           pW8  pV^8”  dB   \        P                  ! W) RR7      pVRV) R13,          p\        P                  ! VVR	RR
7      p\        P                  ! V\        R4      ) V 4      # )a7  
Apply min-p sampling to the logprobs.

Min-p keeps all tokens that are above a minimum probability, scaled by the
probability of the most likely token. As a result, the filter is more
aggressive given a very high-probability token.

Args:
    logprobs: A vector of log probabilities.
    min_p (float): Minimum token probability. Typical values are in the
        0.01-0.2 range, comparably selective as setting `top_p` in the
        0.99-0.8 range.
    min_tokens_to_keep (int, optional): Minimum number of tokens that cannot
        be filtered. Default: ``1``.

r&   z9`min_p` has to be a float in the [0, 1] interval, but is z:`min_tokens_to_keep` has to be a positive integer, but is T)r    Úkeepdimsr]   .NFr   r_   r!   )rb   ra   r   r   ÚmaxÚmathÚlogrc   rd   Úwherer   )r3   r   r   Útop_logprobsÚscaled_min_pÚtokens_to_removeÚtop_indicess   &&&    r   r,   r,   š   sô   € ð, �Ö˜#ÖÜØGÈÀwÐOó
ð 	
ô Ð(¬#×.Ò.Ð3EÈÔ3IÜØHÐI[ÐH\Ð]ó
ð 	
ô
 —6’6˜(¨°dÔ;€LØ¤$§(¢(¨5£/Õ1€LØÑ.Ðð ˜AÔÜ—o’o hÐ4GÈbÔQˆØ! #Ð(:Ð':Ñ';Ð";Õ<ˆÜ×,Ò,ØØØØô	
Ðô �8Š8Ð$¤u¨U£| m°XÓ>Ð>r   c                ód   € V ^8„  d   QhR\         P                  R\        R\         P                  /# )r
   r3   r   r   )r   r   r   )r   s   "r   r   r   Í   s)   € ÷  ñ  œ"Ÿ(™(ð  ¬5ð  ´R·X±Xñ  r   c           	     ó  € \         P                  ! V 4      p\         P                  ! V RR7      p\         P                  ! W#RR7      p\         P                  ! VRR7      p\         P
                  ! \         P                  ! V4      V\         P                  ! VP                  R,          VP                  R7      RR7      p\         P                  ! WVRR7      p\         P                  ! V^V,
          8„  V \        R4      ) 4      # )zÞ
Apply top-p (nucleus) sampling to logits.

Args:
    logprobs: A vector of log probabilities.
    top_p: The cumulative probability threshold for top-p filtering.
Returns:
    token selected based on the top-p criterion.
r   )re   r_   r!   )r   ÚexpÚargsortÚtake_along_axisÚcumsumrd   Ú
zeros_likeÚaranger`   re   ro   r   )r3   r   ÚprobsÚsorted_indicesÚsorted_probsÚcumulative_probsÚinverse_indicess   &&     r   r)   r)   Ì   sÊ   € ô �FŠF�8Ó€Eä—Z’Z ¨rÔ2€NÜ×%Ò% eÀ"ÔE€Lä—y’y °BÔ7Ðô ×'Ò'Ü
�Š�nÓ%ØÜ
�	Š	�.×&Ñ& rÕ*°.×2FÑ2FÔGØô	€Oô ×)Ò)Ð*:ÐRTÔUÐô �8Š8Ø˜1˜u�9Ñ$ØÜ	ˆu‹ˆóð r   c          
      ó’   € V ^8„  d   QhR\         P                  R\        R\        R\        \        ,          R\         P                  /# )r
   rH   r   r   r   r   )r   r   r   r   r   )r   s   "r   r   r   ñ   sF   € ÷ !ñ !Ü�H‰Hð!äð!ô ð!ô œS�	ð	!ô
 ‡X�Xñ!r   c           	     óø  € ^ Tu;8:  d   R8:  g   M \        RV 24      h^ Tu;8:  d   R8:  g   M \        RV 24      h\        P                  ! V R4      pV\        P                  ! WB8„  V\        P                  4      P                  4       8„  pV'       d   RVRV3&   \        P                  ! \        P                  P                  ^ ^4      V8„  V \        P                  ! V\        P                  ) V 4      4      # )aa  
Apply XTC sampling to the logits.

Args:
    logits: The logits from the model's output.
    xtc_probability (float): Probability of XTC sampling to happen for each token
    xtc_threshold (float): The threshold the probs need to reach for being sampled.
    special_tokens_ids (list(int)): List of special tokens IDs to be excluded from XTC sampling.
g      à?z?`threshold` has to be a float in the [0, 0.5] interval, but is r&   z?`probability` has to be a float in the [0, 1] interval, but is F.r!   )rb   r   Úsoftmaxro   r_   ÚminÚrandomÚuniform)rH   r   r   r   r|   Úmasks   &&&&  r   r.   r.   ð   sÞ   € ð  �Ö% #Ö%ÜØMÈmÈ_Ð]ó
ð 	
ð �Ö' CÖ'ÜØMÈoÐM^Ð_ó
ð 	
ô �JŠJ�v˜rÓ"€EØ”2—8’8˜EÑ1°5¼"¿&¹&ÓA×EÑEÓGÑG€DßØ(-ˆˆSÐ$Ð$Ñ%ä�8Š8Ü
�	‰	×Ñ˜!˜QÓ /Ñ1ØÜ
�Š�œŸ™�w Ó'óð r   c                 ó\   € \         P                  P                  V ^V,          ,          4      # ©r   )r   r…   Úcategorical)rH   r   s   &&r   r2   r2     s    € ä�9‰9× Ñ  ¨1¨t­8Õ!4Ó5Ð5r   c                ó0   € V ^8„  d   QhR\         R\        /# ©r
   rU   rV   ©r   r   )r   s   "r   r   r     s   € ÷ (ñ (¤Uð (¼#ñ (r   c                óz   a a€ S ^ 8  g   \        S \        \        34      '       g   \        RS  24      hVV 3R lpV# )aP  
Make repetition penalty processor.

Paper: https://arxiv.org/abs/1909.05858

Args:
    penalty (float): The repetition penalty factor to be applied.
    context_size (int): The number of previous tokens to use.
        Default: ``20``.

Returns:
    Callable[[mx.array, List[int]], mx.array]:
        The repetition penalty processor.
z*penalty must be a non-negative float, got c                 ó¨   <€ \        V 4      ^ 8”  dA   V S) R p VRV 3,          p\        P                  ! V^ 8  VS,          VS,          4      pW!RV 3&   V# ©r   NrD   )Úlenr   ro   )ÚtokensrH   Úselected_logitsrV   rU   s   && €€r   Úrepetition_penalty_processorÚ=make_repetition_penalty.<locals>.repetition_penalty_processor,  sc   ø€ Üˆv‹;˜Œ?Ø˜\˜M˜NÐ+ˆFØ$ Q¨ YÕ/ˆOÜ ŸhšhØ !Ñ#Ø 'Õ)Ø 'Õ)óˆOð
 !0�1�f�9ÑØˆr   )ra   r   r   rb   )rU   rV   r”   s   ff r   rO   rO     s<   ù€ ð �„{œ* W¬s´E¨l×;Ò;ÜÐEÀgÀYÐOÓPÐPö
ð (Ð'r   c                ó0   € V ^8„  d   QhR\         R\        /# rŒ   r�   )r   s   "r   r   r   ;  s   € ÷ &ñ &¤5ð &¼ñ &r   c                ó   a a€ VV 3R lpV# )a¤  
Make a presence penalty processor.

Corresponds to the OpenAI option with the same name. Namely, subtracts
``penalty`` from a logit if the token has occured at least once in the
``context_size`` previous tokens.

Args:
    penalty (float): The presence penalty to be applied.
    context_size (int): The number of previous tokens to use.
        Default: ``20``.

Returns:
    Callable[[mx.array, List[int]], mx.array]
c                 ób   <€ \        V 4      ^ 8”  d   V S) R p VRV 3;;,          S,          uu&   V# r�   )r‘   ©r’   rH   rV   rU   s   &&€€r   Úpresence_penalty_processorÚ9make_presence_penalty.<locals>.presence_penalty_processorL  s5   ø€ Üˆv‹;˜Œ?Ø˜\˜M˜NÐ+ˆFØ�1�f�9× Õ(ÓØˆr   © )rU   rV   rš   s   ff r   rP   rP   ;  s   ù€ ö"ð &Ð%r   c                ó0   € V ^8„  d   QhR\         R\        /# rŒ   r�   )r   s   "r   r   r   U  s   € ÷ 'ñ '¤Eð '¼ñ 'r   c                ó   a a€ VV 3R lpV# )a  
Make a frequency penalty processor.

Corresponds to the OpenAI option with the same name. Namely, subtracts
``penalty`` from a logit for every time that the token has occured in the
``context_size`` previous tokens.

The difference with the presence penalty is that the more often a token
occurs the more it will be penalized.

Args:
    penalty (float): The frequency penalty to be applied.
    context_size (int): The number of previous tokens to use.
        Default: ``20``.

Returns:
    Callable[[mx.array, List[int]], mx.array]
c                 ó|   <€ \        V 4      ^ 8”  d+   V S) R p VP                  RV 3,          P                  S4      pV# r�   )r‘   rE   Úsubtractr™   s   &&€€r   Úfrequency_penalty_processorÚ;make_frequency_penalty.<locals>.frequency_penalty_processori  s>   ø€ Üˆv‹;˜Œ?Ø˜\˜M˜NÐ+ˆFØ—Y‘Y˜q &˜yÕ)×2Ñ2°7Ó;ˆFØˆr   rœ   )rU   rV   r¡   s   ff r   rQ   rQ   U  s   ù€ ö(ð 'Ð&r   )NNé   Nr£   Nr£   r‰   )r£   )rm   Ú	functoolsr   Útypingr   r   r   r   Úmlx.coreÚcorer   r9   rW   Úcompiler…   Ústater0   r,   r)   r.   r2   rO   rP   rQ   rœ   r   r   Ú<module>rª      sN  ðó Ý ß 1Ó 1å ð ØØØØØ ØØ$&÷;÷|6ñr 	ˆ�‰˜BŸI™IŸO™O°R·Y±Y·_±_ÔEôó Fðñ0 	ˆ�‰˜BŸI™IŸO™O°R·Y±Y·_±_ÔEö.?ó Fð.?ñb 	ˆ�‰˜BŸI™IŸO™O°R·Y±Y·_±_ÔEô ó Fð ñF 	ˆ�‰˜BŸI™IŸO™O°R·Y±Y·_±_ÔEô!ó Fð!ñH 	ˆ�‰˜BŸI™IŸO™O°R·Y±Y·_±_ÔEñ6ó Fð6÷(÷B&÷4'ñ 'r   