
    i)B              -          d dl Z d dlZd dlmZmZmZmZmZ d dlZ	d dl
m
Z
  e j                  e      ZddhZg dZg dZg dg d	g d
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fdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d(de	j.                  de	j.                  deedf   deeeee   f      dededee   dee   dedee	j.                     dedededee	j.                     dee	j.                     d ed!ed"ed#ed$edeeef   f*d%Zy))    N)DictListOptionalTupleUnion)tqdmeulerheun)      ?g       @      @)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                     d}|t        |      }|r#|d   dk(  r|j                          |r	|d   dk(  rt        |      dk  rt        j	                  d        nUt        |      dkD  r$t        j	                  dt        |             |dd }|D cg c]  }t        t        |fd	       }}|}|J|H|dkD  rCt        |      D cg c]
  }d
||z  z
   }} d
k7  r|D cg c]  } |z  d
 d
z
  |z  z   z   }}|}|; }	t        t         fd	       |	 k7  rt        j	                  d|	        t            }|S c c}w c c}w c c}w )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                     t        | |z
        S Nabs)xts     V/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/models/mlx/dit_generate.py<lambda>z'get_timestep_schedule.<locals>.<lambda>S   s    c!a%j    keyr   c                      t        | z
        S r)   r*   )r,   r    s    r.   r/   z'get_timestep_schedule.<locals>.<lambda>^   s    AIr0   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_schedulerD   2   sh   * 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   D;E  Eguidance_scalemomentum_statenorm_thresholdc                    ddl m} d}| |z
  }|||j                  dd      z   }||d<   |dkD  rQ|j                  ||z  j	                  |d            }|j                  |j                  |      ||dz   z        }	||	z  }| |j                  | | z  j	                  |d            dz   z  }
||
z  j	                  |d      |
z  }||z
  }| |dz
  |z  z   S )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_uncondrE   rF   rG   mx	proj_axisdiff	diff_normscale_factorv1parallel
orthogonals                r.   _mlx_apg_forwardr^   f   s     I{"D!n((A66$(y!GGTD[--9t-LM	zz",,y"9>YY]M];^_l"	bggy9499ySW9XY\``	aBr	I=BHJ*j888r0   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compile_modeldisable_tqdmsampler_modevelocity_norm_thresholdvelocity_ema_factorc                    	CDEFGHIJ ddl mG ddlm} |t        vrt        d| dt         d      |dk(  IdkD  JdkD  HIr0|d	k(  rt        j                  d
       nt        j                  d       Jrt        j                  d       Hrt        j                  d       i }t        j                         }Gj                  |      }Gj                  |      }|Gj                  |      nd}|Gj                  |      nd}|d   D|d   }|d   }	dkD  xr |
duEErGj                  |
      nd}ErGj                  ||j                        }Gj                  ||gd      }Gj                  ||gd      }|1Gj                  ||j                        } Gj                  || gd      }|Gj                  ||gd      }Eri ndF|Gj                  j                  D||f      }!nt!        |t"              rg }"|D ]  }#|#|#dk  r.|"j%                  Gj                  j                  d||f             8Gj                  j'                  t)        |#            }$|"j%                  Gj                  j                  d||f|$              Gj                  |"d      }!nDGj                  j'                  t)        |            }$Gj                  j                  D||f|$      }!t+        |||      }%t-        |%      }&t)        |&|z        }'dC|r, fd}(	 Gj/                  |(      Ct        j                  d       Ird}*nC |       nd}*|!}+d},CE GIfd}-DE	Ffd}.GHJfd}/t        j                         }0d}1t3        t5        |&      d|      D ]  }2|%|2   }3|2|'k\  r|1sd}1||}|}|* |       }*ErGj                  |+|+gd      n|+}4 |-|4|3|||*      \  }5}*Gj7                  |5        |.|5|3      }5 |/|5|+|,      }5|2|&dz
  k(  r0Gj9                  Dddf|3      }6|+|5|6z  z
  }+Gj7                  |+       n*|%|2dz      }7Ir|dk(  r|3|7z
  }8Gj9                  Dddf|8      }9|+|5|9z  z
  }:Gj7                  |:       ErGj                  |:|:gd      n|:}; |-|;|7|||*      \  }<}*Gj7                  |<        |.|<|7      }< |/|<|:|5      }<d|5|<z   z  }=|+|=|9z  z
  }+|=}5nx|d	k(  rQGj9                  Dddf|3      }6|+|5|6z  z
  }>Gj                  j                  |+j                        }?|7|?z  d|7z
  |>z  z   }+n"|3|7z
  }8Gj9                  Dddf|8      }9|+|5|9z  z
  }+Gj7                  |+       |5}, t        j                         }@t        j                         }A|@|0z
  |d<   |d   t;        |&d      z  |d <   |A|z
  |d!<   ||d"<   t=        j                  |+      }B|B|d#S # t0        $ r!})t        j                  d|)       Y d})~)d})~)ww xY w)$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).   r   )rK   r1   )r"   c           	      ,     | ||||d d      \  }}|S )NFhidden_statestimestep
timestep_rencoder_hidden_statescontext_latentscache	use_cache )xtr-   trencctxvt_mlx_decoders          r.   	_raw_stepz)mlx_generate_diffusion.<locals>._raw_step  s)     1&)3eEB
 Ir0   z4[MLX-DiT] Diffusion step compiled with mx.compile().z:[MLX-DiT] mx.compile() failed (%s); using uncompiled path.c           
          
j                  | j                  d   f|      } | ||||      |fS  	| ||||| xr        \  }}||fS )zSingle model evaluation helper.r   ru   )fullshape)x_inputt_valr   ctx_in
step_cachet_arrvt_out_compiled_stepdo_cfgr   rV   use_heuns          r.   _model_evalz+mlx_generate_diffusion.<locals>._model_eval(  sw    q)+U3%!'5%fEzQQ(!"%"!z2(l

 z!!r0   c                 ^    s| S | d }| d }|cxk  rk  rn |S t        ||	      S |S )zApply CFG guidance if enabled.N)r^   )
vt_rawcurrent_t_valrT   rU   bszrf   re   r   rE   rF   s
       r.   
_apply_cfgz*mlx_generate_diffusion.<locals>._apply_cfg8  sN    M4CL	STlB2BB $I{NN[[r0   c                 (   r|j                  | | z  j                  dd            }j                  ||z  j                  dd            dz   }j                  j                  |      
|z  |dz   z        }| |z  } r|d	z
  | z  	|z  z   } | S )z/Apply optional norm clamping and EMA smoothing.)r&   rs   TrJ   g|=r   )rP   rQ   rR   rS   )	vt_guided
xt_currentprev_velocityvt_normxt_normscalerV   use_emause_norm_clamprn   rm   s         r.   _apply_stabilisationz4mlx_generate_diffusion.<locals>._apply_stabilisationB  s     ggy9499vPT9UVGggzJ6;;RV;WX[``GJJW%(72wGE "E)I }022i?BUXeBeeIr0   FzMLX DiT diffusion)descdisableToder   diffusion_time_costdiffusion_per_step_time_costtotal_time_costrl   )target_latents
time_costs)rM   rN   	dit_modelrp   VALID_SAMPLER_MODES
ValueErrorr7   r8   infotimearraybroadcast_tor   concatenaterandomnormal
isinstancer4   appendr2   intrD   r6   compile	Exceptionr   r;   evalr   maxnp)Kr   r_   r`   ra   rb   rc   r    r!   r"   rE   rd   re   rf   rg   rh   ri   rj   rk   rl   rm   rn   rp   r   total_startenc_hsr   	enc_hs_ncctx_ncTC	null_condnull_expandednull_expanded_ncnoisepartssr2   r>   	num_stepscover_stepsr   excr{   r~   prev_vtr   r   r   
diff_start_switched_to_non_coverstep_idx	current_tx_inr   t_unsqnext_tdtdt_arrxt_predictedx_in2vt2vt_avg
pred_clean	new_noisediff_end	total_end	result_npr   r   r   rF   rV   r   r   r   sK   `        ` ``      ``                                              @@@@@@@@r.   mlx_generate_diffusionr      s   x 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
A
C!A!A c!G&;4&GF39./tI	6<<@ 7a@nnc3Zan0 !y)//J	3C'D1MI^^VV$41^=F+1RtN |		  #q!-	D$	AyAERYY--q!Qi89iimmCF+RYY--q!QiS-AB  u1-iimmCI&		  #q!# 6 ,E9+VOO$Ii"667K N		ZZ	2NKKNO ,:,B&(	BG" "   $ J"y)0C\Z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} [@ yy{H		I(0:(=J$%1;<Q1RUXYbdeUf1fJ-.$-$;J !!-J~I#  O  	NNLc 	s   &W! !	X*XX)r   NN)Ng      @)Nr   r   NNr   N        r   r   NNFFr	   r   r   )loggingr   typingr   r   r   r   r   numpyr   r   	getLogger__name__r7   r   r<   r:   r=   floatr4   r   rD   r^   ndarraystrboolobjectr   r}   r0   r.   <module>r      s:  (   5 5  			8	$'   
<
K
(	  $!%11~1 #1 
%[	1p &*9 9 TN	9
 9L -1 $!%26 #!"%?C9=%(!$+\ jj\ 

\ S#X	\
 5d3i(
)\ \ \ ~\ #\ \ $BJJ/\ \ \  \ )1(<\  #+2::"6!\" #\$ %\& '\( #)\* +\, 
#v+-\r0   