
    hti                          S SK r S SKJs  Jr   " S S5      rS/4S\ R                  S\ R                  4S jjrSSS	S/4S
\ R                  S\ R                  S\S\S\S\4S jjr	S r
S rS r   SS\ R                  S\ R                  S\ R                  S\ R                  S\S\S\S\4S jjr SS\ R                  S\ R                  S\ R                  S\S\S\4S jjrS\ R                  S\ R                  S\ R                  S\S\4
S jrg)     Nc                   H    \ rS rSrSS\4S jjrS\R                  4S jrSr	g)	MomentumBuffer   momentumc                     Xl         SU l        g )Nr   r   running_average)selfr   s     Y/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/models/common/apg_guidance.py__init__MomentumBuffer.__init__   s          update_valuec                 H    U R                   U R                  -  nX-   U l        g Nr   )r
   r   new_averages      r   updateMomentumBuffer.update   s!    mmd&:&::+9r   r   N)g      )
__name__
__module____qualname____firstlineno__floatr   torchTensorr   __static_attributes__ r   r   r   r      s    ! !:5<< :r   r   v0v1c                    U R                   nU R                  R                  nUS:X  a  U R                  5       UR                  5       pU R	                  5       UR	                  5       p[
        R                  R                  R                  XS9nX-  R                  USS9U-  nX-
  nUR                  U5      R                  U5      UR                  U5      R                  U5      4$ )NmpsdimTr$   keepdim)dtypedevicetypecpudoubler   nn
functional	normalizesumto)r   r    dimsr'   device_typev0_parallelv0_orthogonals          r   projectr5      s    
 HHE))..Ke2668BYY["))+				&	&r	&	4B7--D$-7"<K$M>>% ##K0-2B2B52I2L2L[2YYYr   g        g      @	pred_condpred_uncondguidance_scalemomentum_bufferetanorm_thresholdc                     X-
  nUb  UR                  U5        UR                  nUS:  aD  [        R                  " U5      nUR	                  SUSS9n	[        R
                  " XU	-  5      n
Xz-  n[        XpU5      u  pXU-  -   nXS-
  U-  -   nU$ )Nr      T)pr$   r&      )r   r	   r   	ones_likenormminimumr5   )r6   r7   r8   r9   r:   r;   r1   diffones	diff_normscale_factordiff_paralleldiff_orthogonalnormalized_updatepred_guideds                  r   apg_forwardrK   !   s     "D"t$..t$IItTI:	}}TI+EF"%,Td%C"M'*==15FFFKr   c                     XX-
  -  -   $ r   r   )cond_outputuncond_outputcfg_strengths      r   cfg_forwardrP   ;   s    ;+FGGGr   c                     U [         R                  R                  U SSS9-  n U[         R                  R                  USSS9-  n[         R                  " X-  SSS9nU$ )z
Calculate cosine similarity between two normalized tensors.

Args:
    tensor1: First tensor [B, ...]
    tensor2: Second tensor [B, ...]

Returns:
    Cosine similarity value [B, 1]
r?   Tr%   )r   linalgrA   r/   )tensor1tensor2cosvalues      r   call_cos_tensorrV   ?   s^     ))'q$)GGG))'q$)GGGyy*4@HOr   c                    U R                   u  p#nU R                  X#-  U5      R                  5       n UR                  X#-  U5      R                  5       nU R                  5       UR                  5       :w  a  [	        S5      e[
        R                  " X-  SSS9n[
        R                  " X-  SSS9nXVS-   -  U-  nX-
  nUR                  X#U5      UR                  X#U5      4$ )a<  
Decompose latent_diff into parallel and perpendicular components relative to latent_hat_uncond.

Args:
    latent_diff: Difference tensor [B, C, ...]
    latent_hat_uncond: Unconditional prediction tensor [B, C, ...]

Returns:
    projection: Parallel component
    perpendicular_component: Perpendicular component
zBlatent_diff and latent_hat_uncond must have the same shape [n, d].r?   Tr%   g:0yE>)shapeviewr   size
ValueErrorr   r/   reshape)	latent_difflatent_hat_uncondntcdot_productnorm_square
projectionperpendicular_components	            r   compute_perpendicular_componentrf   P   s     GA!""15!,224K)..qua8>>@.3355]^^))K;DQK))-AqRVWKt!348IIJ)6??1#%<%D%DQ1%MMMr   latentsnoise_pred_condnoise_pred_uncondsigma
angle_clip
apply_norm
apply_clipc                 R   U R                   S   UR                   S   :w  a]  UR                   S   U R                   S   -  S:w  a  [        S5      eUR                   S   U R                   S   -  nU R                  USS9n UR                   S   n	Un
U
R                   u  pn[        U[        [
        45      (       aN  [        R                  " X0R                  U R                  S9nUR                  SSS5      R                  U	SS5      nO[        R                  " U5      (       a|  UR                  5       S:X  a%  UR                  SSS5      R                  U	SS5      nOZUR                  5       U	:X  a  UR                  U	SS5      nO2[        SU	 SUR                    35      e[        S[        U5       35      eUS-
  nXS:  -  S	-   nXU
-  -
  nXU-  -
  nX-
  n[!        UR                  S
U5      R#                  [
        5      UR%                  S
U5      R'                  5       R#                  [
        5      5      R)                  SS5      n[        R*                  " U5      R                  XS5      nU(       a  [        R,                  " UU-  U* U5      OUU-  n[/        UU5      u  nn[        R0                  " U5      U-  nU[        R2                  " U5      -  [        R2                  " U5      -  [        R2                  " U5      S	:  -  UU-  [        R2                  " U5      S	:*  -  -   nUU-   nU(       aB  U[        R4                  R7                  USSS9-  [        R4                  R7                  USSS9-  nU U-
  U-  nUR%                  XU5      R#                  U R                  5      nU$ )a  
ADG (Angle-based Dynamic Guidance) forward pass for Flow Matching.

In flow matching (including SD3), sigma represents the current timestep t_curr.
The predictions are velocity fields v(x_t, t).

Args:
    latents: Current state x_t [N, T, d] where d=64
    noise_pred_cond: Conditional velocity prediction v_cond [N, T, d]
    noise_pred_uncond: Unconditional velocity prediction v_uncond [N, T, d]
    sigma: Current timestep t_curr (not t_prev!)
    guidance_scale: Guidance strength
    angle_clip: Maximum angle for clipping (default: pi/6)
    apply_norm: Whether to normalize the result (ADG_w_norm variant)
    apply_clip: Whether to clip the angle (ADG_wo_clip when False)

Returns:
    Guided velocity prediction [N, T, d]
r?   r   zYnoise_pred_cond time dimension must be a whole-number multiple of latents time dimension.r#   )r(   r'   z6sigma has incompatible shape. Expected scalar or size z, got z&sigma must be a number or tensor, got gMbP?r   g!g!?Tr%   )rX   r[   repeat_interleave
isinstanceintr   r   tensorr(   r'   rY   expand	is_tensornumel	TypeErrorr)   rV   r0   r\   
contiguousclampacoscliprf   cossinrR   rA   )rg   rh   ri   rj   r8   rk   rl   rm   repeatsr_   noise_pred_textr`   ra   weightlatent_hat_textr^   r]   	cos_thetalatent_thetalatent_theta_newprojperplatent_v_newlatent_p_new
latent_new
noise_preds                             r   adg_forwardr   k   sv   : }}Q?0033  #gmmA&66!;xyy!''*gmmA.>>++G+; 	a A%O##GA! %#u&&U>>O

1a#**1a3			;;=AJJq!Q'..q!Q7E[[]aJJq!Q'EUVWUXX^_d_j_j^klmm@eNOO aFz"T)F 77O*;";;!5K  R#&&u-!!"a(33588? eK$  ::i(--aA6LU_uzz&<"7*jQeknzez0>OPJD$99-.@L%))$455		,8OO		,$&(*.-599\;RVZ;Z*[\L,J%,,"3"3OTX"3"YY\a\h\h\m\mAt ]n ]- -
 J&%/J##A!,//>Jr   c                 "    [        U UUUUUSSS9$ )z
ADG with normalization - preserves the magnitude of latent predictions.

This variant normalizes the final latent to maintain the same norm as the
conditional prediction, which can help preserve image quality.
T)rk   rl   rm   r   )rg   rh   ri   rj   r8   rk   s         r   adg_w_norm_forwardr      s(     w&(%","&"&( (r   c           
          [        XX#USSS9$ )z
ADG without angle clipping - allows unbounded angle adjustments.

This variant doesn't clip the angle, which may result in more aggressive
guidance but could be less stable.
F)rl   rm   r   )rg   rh   ri   rj   r8   s        r   adg_wo_clip_forwardr      s     w1B>fkx}~~r   )gX%?FT)gX%?)r   torch.nn.functionalr,   r-   Fr   r   r5   r   rK   rP   rV   rf   boolr   r   r   r   r   r   <module>r      s     : : ZZZ* '+
||  $	
 
 4H"NB !Q\\Q\\Q ||Q <<	Q
 Q Q Q Qt !(\\(\\( ||( 	(
 ( (0\\\\ || 	
 r   