
    0	j/U              C          S SK r S SKrS SKrS SKJrJrJrJrJr  S SK	r
S SKJr  \ R                  " \5      rSS1r/ SQr/ SQr/ SQ/ S	Q/ S
QS.r   S3S\S\\   S\\   S\\   4S jjr  S4S\S\\   S\4S jjrS rS r                            S5S\
R4                  S\
R4                  S\\S4   S\\\\\   4      S\S\S\\   S\\   S\S\\
R4                     S\S\S\S \\
R4                     S!\\
R4                     S"\\\\\   4      S#\S$\S%\S&\S'\S(\S)\S*\S+\S,\S-\S.\\
R4                     S/\\
R4                     S0\S1\S\\\4   4@S2 jjrg)6    N)DictListOptionalTupleUnion)tqdmeulerheun)      ?       @      @)r   .袋???      ?۶m۶m?竪?ى؉?      ?UUUUUU?%I$I?      ?tE]t?      ?皙?      ?333333?      ?qq?      ?)r   r   r   r   r   r   r   r    )r   r   r   r   r   r   r   r   )r   r   r   r   r   r   r   r   shift	timestepsinfer_stepsreturnc                   ^  SnUb  [        U5      nU(       a+  US   S:X  a"  UR                  5         U(       a  US   S:X  a  M"  [        U5      S:  a  [        R	                  ST 5        OW[        U5      S:  a$  [        R	                  S[        U5      5        USS nU Vs/ s H  n[        [        U4S jS	9PM     nnUnUcR  UbO  US:  aI  [        U5       Vs/ s H
  nS
Xr-  -
  PM     nnT S
:w  a!  U Vs/ s H  nT U-  S
T S
-
  U-  -   -  PM     nnUnUc:  T n	[        [        U 4S jS	9m U	T :w  a  [        R	                  SU	T 5        [        T    nU$ s  snf s  snf s  snf )a  Compute the timestep schedule for diffusion sampling.

When ``infer_steps`` is provided and ``timesteps`` is None, a continuous
linspace schedule is generated (matching the PyTorch base-model behaviour).
The legacy lookup-table path (8-step ``SHIFT_TIMESTEPS``) is used only when
neither ``timesteps`` nor ``infer_steps`` is supplied.

Args:
    shift: Diffusion timestep shift (applied via ``shift*t / (1+(shift-1)*t)``).
    timesteps: Optional custom list of timesteps.
    infer_steps: Number of diffusion steps.  When given, overrides the
        fixed 8-step lookup table.

Returns:
    List of timestep values (descending, without trailing 0).
Nr      z<timesteps empty after removing zeros; using default shift=%s   z$timesteps length=%d > 20; truncatingc                     [        X-
  5      $ Nabs)xts     V/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/models/mlx/dit_generate.py<lambda>'get_timestep_schedule.<locals>.<lambda>T   s
    c!%j    keyr   c                     > [        U T-
  5      $ r*   r+   )r-   r!   s    r/   r0   r1   _   s    AIr2   z.shift=%.2f rounded to nearest valid shift=%.1f)
listpoplenloggerwarningminVALID_TIMESTEPSrangeVALID_SHIFTSSHIFT_TIMESTEPS)
r!   r"   r#   t_schedule_listts_listr.   mappediraworiginal_shifts
   `         r/   get_timestep_schedulerF   3   sf   * Oy/'"+*KKM '"+*w<!NNY[`a7|b Es7|T!#2,SZ[SZac/1/HISZF[$O;#:{Q.3K.@A.@sQ_$.@AC<BEF#Q519us{a&7 78#CFL&>?U"NNK^]bc)%0! \ BFs   !EE7Eguidance_scalemomentum_statenorm_thresholdc                    SSK Jn  SnX-
  nUb  XsR                  SS5      -   nXsS'   US:  aK  UR                  Xw-  R	                  USS95      nUR                  UR                  U5      XHS-   -  5      n	Xy-  nXR                  X -  R	                  USS95      S-   -  n
Xz-  R	                  USS9U
-  nX{-
  nXS-
  U-  -   $ )u   APG (Adaptive Projected Guidance) in pure MLX — mirrors the PyTorch ``apg_forward``.

Projection is performed along axis 1 (the time/sequence dimension) to match
the PyTorch implementation which calls ``apg_forward(..., dims=[1])``.
r   Nr'   runningTaxiskeepdimsg:0yE>)mlx.corecoregetsqrtsumminimum	ones_like)	pred_condpred_uncondrG   rH   rI   mx	proj_axisdiff	diff_normscale_factorv1parallel
orthogonals                r/   _mlx_apg_forwardr`   g   s     I"D!((A66$(y!GGT[--9t-LM	zz",,y"9>Y]M];^_"	ggy499ySW9XY\``	aB	I=BHJ*j888r2   c                 l    SSK Jn  X4-  SU-
  U-  -   nUR                  USS9nUR                  XpU5      $ )zEReplace non-repaint regions of *xt* with noised source latents (MLX).r   Nr   r&   rM   )rO   rP   expand_dimswhere)xt	clean_srcmaskt_nextnoiserX   ztms           r/   _mlx_repaint_step_injectionrl      s=    	3<94	4B
t"%A88A2r2   c                 D   SSK Jn  UR                  [        R                  5      R                  5       nUS::  a,  UR                  UR                  U5      SS9nX`-  SU-
  U-  -   $ UR                  u  px[        U5       H  n	X)   n
U
R                  5       (       d  U
R                  5       (       d  M3  [        R                  " U
5      S   n[        U5      S:X  a  M]  [        US   5      [        US   5      S-   p[        X-
  S5      nX-
  S:  a%  [        R                   " SSX-
  S-   5      SS XYX24'   [#        X-   U5      nX-
  S:  d  M  [        R                   " SSX-
  S-   5      SS XYX24'   M     UR                  UR                  U5      SS9nX`-  SU-
  U-  -   $ )z@Blend generated latents with source at repaint boundaries (MLX).r   Nr&   rb   r   r'      )rO   rP   astypenpfloat32copyrc   arrayshaper=   allanynonzeror8   intmaxlinspacer;   )x_genrf   mask_np	cf_framesrX   softrk   BTbrowidxleftrightfsfes                   r/   _mlx_repaint_boundary_blendr      s}   >>"**%**,DA~NN288D>N3yC!Gy000==DA1Xj7799CGGIIjjoa s8q=#a&k3s2w<!#3e!1%9q=!{{1aQ?"EDBG"A&:> "Aq"*q. A!B GDEH  	rxx~B/A9a9,,,r2   encoder_hidden_states_npcontext_latents_npsrc_latents_shape.seedinfer_methodnull_condition_emb_npcfg_interval_startcfg_interval_endaudio_cover_strength"encoder_hidden_states_non_cover_npcontext_latents_non_cover_npretake_seedretake_variancecompile_modeldisable_tqdmsampler_modevelocity_norm_thresholdvelocity_ema_factordcw_enableddcw_mode
dcw_scalerdcw_high_scalerdcw_waveletrepaint_mask_npclean_src_latents_nprepaint_crossfade_framesrepaint_injection_ratioc                    ^ ^	^^^^^X^Y^Z^[^\^]^^^_^`^a SSK Jm^  SSKJn   U[        ;  a  [        SU S[         S35      eUS:H  m`TS:  maTS:  m_T`(       a1  US	:X  a  [        R                  S
5        O[        R                  S5        Ta(       a  [        R                  ST5        T_(       a  [        R                  ST5        0 n![        R                  " 5       n"T^R                  U5      n#T^R                  U5      n$Ub  T^R                  U5      OSn%Ub  T^R                  U5      OSn&USL=(       a    USLn'U'(       a  T^R                  U5      OSn(U'(       a  T^R                  U5      OSn)US   m[US   mYUS   mXT	S:  =(       a    U
SLm\T\(       a  T^R                  U
5      OSn*T\(       a  T^R                  U*U#R                  5      n+T^R                  U#U+/SS9n#T^R                  U$U$/SS9n$U%b.  T^R                  U*U%R                  5      n,T^R                  U%U,/SS9n%U&b  T^R                  U&U&/SS9n&T\(       a  0 OSm]UXUYU[U^4S jn-U-" U5      n.US:  aQ  U-" U5      n/U[        R                  S-  -  n0[        R                   " U05      U.-  [        R"                  " U05      U/-  -   n.[%        XgUS9n1['        U15      n2[)        U2U-  5      n3SmZU(       a-  U 4S jn4 T^R+                  U45      mZ[        R                  S5        T`(       a  Sn6OTZc  U " 5       OSn6U.n7Sn8UZU\U U^U`4S jn9U[UUU\U	U]4S jn:U^U_UaUU4S jn;[        R                  " 5       n<Sn=SSKJn>  U=(       a    US:g  =(       d    US:H  =(       a    US:g  n?U?(       a(  US:X  a  SOS U S!3n@[        R                  S"UUUUU@5        [3        [5        U25      S#US$9 GH  nAU1UA   nBUAU3:  a  U=(       d  S%n=U%b  U%n#U&n$U6b  U " 5       n6T\(       a  T^R                  U7U7/SS9OU7nCU9" UCWBU#U$U65      u  nDn6T^R7                  UD5        U:" UDUB5      nDU;" UDU7U85      nDU7nEUDnFWAU2S-
  :X  a0  T^R9                  T[SS4WB5      nGU7WDUG-  -
  n7T^R7                  U75        GO3U1WAS-      nHT`(       a  US&:X  a  WBWH-
  nIT^R9                  T[SS4UI5      nJU7WDUJ-  -
  nKT^R7                  UK5        T\(       a  T^R                  WKUK/SS9OWKnLU9" ULWHU#U$U65      u  nMn6T^R7                  UM5        U:" UMUH5      nMU;" UMWKWD5      nMS'UDUM-   -  nNU7UNWJ-  -
  n7UNnDOyUS	:X  aQ  T^R9                  T[SS4WB5      nGU7WDUG-  -
  nOT^R:                  R=                  U7R                  5      nPWHUP-  SUH-
  UO-  -   n7O"WBWH-
  nIT^R9                  T[SS4UI5      nJU7WDUJ-  -
  n7T^R7                  U75        U?(       a;  T^R9                  T[SS4WB5      nQWEWFUQ-  -
  nRU>" U7URUBS%UUUUS(9n7T^R7                  U75        WDn8U'(       d  GM5  [?        UU2-  5      nSWAUS:  d  GML  WAU2S-
  :  a  U1WAS-      OSnT[A        U7U)U(UTU.5      n7T^R7                  U75        GM     U'(       a%  US:  a  [C        U7U)UU5      n7T^R7                  U75        [        R                  " 5       nU[        R                  " 5       nVUUU<-
  U!S)'   U!S)   [E        U2S5      -  U!S*'   UVU"-
  U!S+'   UU!S,'   [F        R                  " U75      nWUWU!S-.$ ! [,         a!  n5[        R                  SU55         Sn5A5GNSn5A5ff = f).u  Run the complete MLX diffusion loop with optional CFG guidance.

This is the core generation function.  It accepts numpy arrays (converted
from PyTorch tensors by the handler) and returns numpy arrays that the
handler converts back to PyTorch.

Args:
    mlx_decoder: ``MLXDiTDecoder`` instance with loaded weights.
    encoder_hidden_states_np: [B, enc_L, D] from prepare_condition (numpy).
    context_latents_np: [B, T, C] from prepare_condition (numpy).
    src_latents_shape: shape tuple [B, T, 64] for noise generation.
    seed: random seed (int, list[int], or None).
    infer_method: "ode" or "sde".
    shift: timestep shift factor.
    timesteps: optional custom timestep list.
    infer_steps: number of diffusion steps.
    guidance_scale: CFG guidance strength (>1.0 enables CFG).
    null_condition_emb_np: [1, 1, D] null condition embedding for CFG.
    cfg_interval_start: timestep ratio below which CFG is disabled.
    cfg_interval_end: timestep ratio above which CFG is disabled.
    audio_cover_strength: cover strength (0-1).
    encoder_hidden_states_non_cover_np: optional [B, enc_L, D] for non-cover.
    context_latents_non_cover_np: optional [B, T, C] for non-cover.
    compile_model: If True, compile the decoder step with ``mx.compile``.
    disable_tqdm: If True, suppress the diffusion progress bar.
    sampler_mode: Sampler algorithm — ``"euler"`` (first-order, default) or
        ``"heun"`` (second-order predictor-corrector for cleaner output).
    velocity_norm_threshold: Clamp velocity prediction L2 norm relative to
        input norm at each step.  0 disables (default).  Values around
        1.5–3.0 reduce outlier artefacts.
    velocity_ema_factor: Blend current velocity prediction with the previous
        step's prediction via EMA (``vt = (1-f)*vt + f*prev``).
        0 disables (default).  Values around 0.05–0.2 smooth the trajectory.

Returns:
    Dict with ``"target_latents"`` (numpy) and ``"time_costs"`` dict.
r   Nr'   )MLXCrossAttentionCachezUnsupported sampler_mode 'z'. Expected one of .r
   sdez[MLX-DiT] Heun sampler is not supported with SDE inference method. Falling back to Euler for SDE steps. Use infer_method='ode' for Heun.zF[MLX-DiT] Using Heun (second-order) sampler for higher-quality output.z:[MLX-DiT] Velocity norm clamping enabled (threshold=%.2f).z7[MLX-DiT] Velocity EMA smoothing enabled (factor=%.3f).rn   r   rb   c                 8  > U c  TR                   R                  TTT45      $ [        U [        5      (       a  / nU  H  nUb  US:  a/  UR	                  TR                   R                  STT45      5        M;  TR                   R                  [        U5      5      nUR	                  TR                   R                  STT4US95        M     TR                  USS9$ TR                   R                  [        U 5      5      nTR                   R                  TTT4US9$ )Nr   r'   r3   rb   )randomnormal
isinstancer6   appendr4   rx   concatenate)_seedpartssr4   Cr   bszrX   s       r/   _draw_noise+mlx_generate_diffusion.<locals>._draw_noise)  s    =99##S!QK00eT""E9ALL!1!11a)!<=))--A/CLL!1!11a)!1!EF  >>%a>00iimmCJ'yya55r2           r   )r#   c           
          > T" XUX4S SS9u  pVU$ )NFhidden_statestimestep
timestep_rencoder_hidden_statescontext_latentscache	use_cache )re   r.   trencctxvt_mlx_decoders          r/   	_raw_step)mlx_generate_diffusion.<locals>._raw_stepI  s#     &)eEB
 Ir2   z4[MLX-DiT] Diffusion step compiled with mx.compile().z:[MLX-DiT] mx.compile() failed (%s); using uncompiled path.c           
         > T
R                  U R                  S   4U5      nTb  T" XXRU5      U4$ T	" U UUUUUT(       + =(       a    T(       + S9u  pdXd4$ )zSingle model evaluation helper.r   r   )fullrt   )x_inputt_valr   ctx_in
step_cachet_arrvt_out_compiled_stepdo_cfgr   rX   use_heuns          r/   _model_eval+mlx_generate_diffusion.<locals>._model_evalc  so    q)+U3%!'%fEzQQ(!"%"!z2(l
 !!r2   c                 j   > T(       d  U $ U ST nU TS nTUs=::  a  T::  a  O  U$ [        X#TT	5      $ U$ )zApply CFG guidance if enabled.N)r`   )
vt_rawcurrent_t_valrV   rW   r   r   r   r   rG   rH   s
       r/   
_apply_cfg*mlx_generate_diffusion.<locals>._apply_cfgs  sL    M4CL	STlB2BB $INN[[r2   c                 ,  > T(       as  TR                  X -  R                  SSS95      nTR                  X-  R                  SSS95      S-   nTR                  TR                  U5      T
U-  US-   -  5      nX-  n T(       a  Ub  ST	-
  U -  T	U-  -   n U $ )z/Apply optional norm clamping and EMA smoothing.)r'   rn   TrL   g|=r   )rR   rS   rT   rU   )	vt_guided
xt_currentprev_velocityvt_normxt_normscalerX   use_emause_norm_clampr   r   s         r/   _apply_stabilisation4mlx_generate_diffusion.<locals>._apply_stabilisation}  s     ggy499vPT9UVGggz6;;RV;WX[``GJJW%(72wGE ")I }022i?BUXeBeeIr2   F)apply_mlx_dcwdoublehaarzMLX-native Haarztorch bridge ()zW[MLX-DiT] DCW enabled (mode=%s, scaler=%.3f, high_scaler=%.3f, wavelet=%s, backend=%s).zMLX DiT diffusion)descdisableToder   )t_currenabledmodescalerhigh_scalerwaveletdiffusion_time_costdiffusion_per_step_time_costtotal_time_costr   )target_latents
time_costs)$rO   rP   	dit_modelr   VALID_SAMPLER_MODES
ValueErrorr9   r:   infotimers   broadcast_tort   r   mathpicossinrF   r8   rx   compile	Exception%acestep.models.mlx.dcw_correction_mlxr   r   r=   evalr   r   r   roundrl   r   ry   rp   )br   r   r   r   r   r   r!   r"   r#   rG   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   total_startenc_hsr   	enc_hs_ncctx_nc
do_repaintrepaint_mask_mxclean_src_mx	null_condnull_expandednull_expanded_ncr   ri   retake_noisev_radr@   	num_stepscover_stepsr   excr   re   prev_vtr   r   r   
diff_start_switched_to_non_coverr   
dcw_active_backendstep_idx	current_tx_inr   xt_before_stepvt_for_denoiset_unsqnext_tdtdt_arrxt_predictedx_in2vt2vt_avg
pred_clean	new_noiset_unsq_ddenoisedinjection_cutofft_afterdiff_end	total_end	result_npr   r   r   r   r   rH   rX   r   r   r   sb   `        ` ``        ``                                                                 @@@@@@@@@@r/   mlx_generate_diffusionr0     s   N 1..5l^CVWjVkklmnnv%H,q0N!A%G5 NNX
 KK`aPRijMObcJ))+KXX./F
((%
&C@b@n;<txI7S7_RXX23eiF !,Q1ET1QJ3=bhh/4O5?28801TL
A
C!A!A c!G&;4&GF39./tI	6<<@ 7a@nnc3Zan0 !y)//J	3C'D1MI^^VV$41^=F+1RtN6 6 E ";/477S=1%'$((5/L*HH ,E+VOO$Ii"667K N		ZZ	2NKKNO ,:,B&(	BG" "   $ J" D cNh(2M#7M  (3v(=$^T_S``aCbej/;	

 y)0C\Z#H-	 {"+A%)"$" .0 4:r~~r2hQ~/r  ieD	E
I&!"b'2  y1}$WWc1a[)4Fb6k!BGGBK$X\2FLE1 '#q!b1!BK/% QWl'C!L\h(UK
U f-*3bA S)&6/)&#q!i8"v+-
II,,RXX6	i'3<:*EE '#q!b1"v+%GGBK
 wwQ{I6H%(AAHHYj+[	B GGBK :$%<y%HI**;CiRSm;S/(Q,7Y\0\?T[]bcy [~ .2(\?Lde
yy{H		I(0:(=J$%1;<Q1RUXYbdeUf1fJ-.$-$;J !!-J~I#  s  	NNLc 	s   &[ 
\ [<<\)r   NN)Ng      @)Nr   r   NNr   Nr   r   r   NNNr   FFr	   r   r   Tr   g?g{Gz?r   NN
   r   )loggingr   r   typingr   r   r   r   r   numpyrp   r   	getLogger__name__r9   r   r>   r<   r?   floatr6   rx   rF   r`   rl   r   ndarraystrboolobjectr0  r   r2   r/   <module>r<     s  (    5 5  			8	$'   
<
K
(	  $!%11~1 #1 
%[	1p &*9 9 TN	9
 9B-> -1 $!%26 #!"%?C9=37 %(!$!,015$&%(Af jjf 

f S#X	f
 5d3i(
)f f f ~f #f f $BJJ/f f f  f )1(<f  #+2::"6!f" %T#Y/0#f$ %f& 'f( )f* +f, #-f. /f0 1f2 3f4 5f6 7f8 9f: bjj);f< #2::.=f> "?f@ #AfB 
#v+Cfr2   