
    hti                         d dl Z d dlmc mZ  G d d      Zdgfde j                  de j                  fdZddd	dgfd
e j                  de j                  dedededefdZ	d Z
d Zd Z	 	 	 dde j                  de j                  de j                  de j                  dedededefdZ	 dde j                  de j                  de j                  dededefdZde j                  de j                  de j                  dedef
dZy)     Nc                   :    e Zd ZddefdZdej                  fdZy)MomentumBuffermomentumc                      || _         d| _        y )Nr   r   running_average)selfr   s     Y/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/models/common/apg_guidance.py__init__zMomentumBuffer.__init__   s          update_valuec                 J    | j                   | j                  z  }||z   | _        y Nr   )r	   r   new_averages      r
   updatezMomentumBuffer.update   s#    mmd&:&::+k9r   N)g      )__name__
__module____qualname__floatr   torchTensorr    r   r
   r   r      s    ! !:5<< :r   r   v0v1c                    | j                   }| j                  j                  }|dk(  r | j                         |j                         }} | j	                         |j	                         }} t
        j                  j                  j                  ||      }| |z  j                  |d      |z  }| |z
  }|j                  |      j                  |      |j                  |      j                  |      fS )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
   projectr0      s    
 HHE))..Ke2668BYY["))+B				&	&rt	&	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                 ,   | |z
  }||j                  |       |j                  }|dkD  rGt        j                  |      }|j	                  d|d      }	t        j
                  |||	z        }
||
z  }t        || |      \  }}|||z  z   }| |dz
  |z  z   }|S )Nr      T)pr   r!      )r   r   r   	ones_likenormminimumr0   )r1   r2   r3   r4   r5   r6   r,   diffones	diff_normscale_factordiff_paralleldiff_orthogonalnormalized_updatepred_guideds                  r
   apg_forwardrF   !   s     {"D"t$..t$IItTI:	}}T>I+EFl"%,T9d%C"M?'#*==~15FFFKr   c                     ||| |z
  z  z   S r   r   )cond_outputuncond_outputcfg_strengths      r
   cfg_forwardrK   ;   s    <;+FGGGr   c                     | t         j                  j                  | dd      z  } |t         j                  j                  |dd      z  }t        j                  | |z  dd      }|S )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   linalgr<   r*   )tensor1tensor2cosvalues      r
   call_cos_tensorrQ   ?   s`     ))'q$)GGG))'q$)GGGyy7*4@HOr   c                    | j                   \  }}}| j                  ||z  |      j                         } |j                  ||z  |      j                         }| j                         |j                         k7  rt	        d      t        j                  | |z  dd      }t        j                  ||z  dd      }||dz   z  |z  }| |z
  }|j                  |||      |j                  |||      fS )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_componentra   P   s     GAq!""1q5!,224K)..q1ua8>>@.3355]^^))K*;;DQK))-0AAqRVWKt!348IIJ)J6??1a#%<%D%DQ1%MMMr   latentsnoise_pred_condnoise_pred_uncondsigma
angle_clip
apply_norm
apply_clipc                 8   | j                   d   |j                   d   k7  r_|j                   d   | j                   d   z  dk7  rt        d      |j                   d   | j                   d   z  }| j                  |d      } |j                   d   }	|}
|
j                   \  }	}}t        |t        t
        f      rQt        j                  || j                  | j                        }|j                  ddd      j                  |	dd      }nt        j                  |      rz|j                         dk(  r%|j                  ddd      j                  |	dd      }nY|j                         |	k(  r|j                  |	dd      }n2t        d|	 d|j                          t        dt        |             |dz
  }||dkD  z  d	z   }| ||
z  z
  }| ||z  z
  }||z
  }t!        |j                  d
|      j#                  t
              |j%                  d
|      j'                         j#                  t
                    j)                  dd      }t        j*                  |      j                  |	|d      }|rt        j,                  ||z  | |      n||z  }t/        ||      \  }}t        j0                  |      |z  }|t        j2                  |      z  t        j2                  |      z  t        j2                  |      d	kD  z  ||z  t        j2                  |      d	k  z  z   }||z   }|rH|t        j4                  j7                  |dd      z  t        j4                  j7                  |dd      z  }| |z
  |z  }|j%                  |	||      j#                  | j                        }|S )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    )rS   rV   repeat_interleave
isinstanceintr   r   tensorr#   r"   rT   expand	is_tensornumel	TypeErrorr$   rQ   r+   rW   
contiguousclampacosclipra   cossinrM   r<   )rb   rc   rd   re   r3   rf   rg   rh   repeatsrZ   noise_pred_textr[   r\   weightlatent_hat_textrY   rX   	cos_thetalatent_thetalatent_theta_newprojperplatent_v_newlatent_p_new
latent_new
noise_preds                             r
   adg_forwardr   k   s   : }}Q?0033  #gmmA&66!;xyy!''*gmmA.>>++G+; 	a A%O##GAq! %#u&U7>>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vz"T)F 77O%*;";;!$55K  R#&&u-!!"a(33588? eK$  ::i(--aA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##Aq!,//>Jr   c           
      (    t        | |||||dd      S )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)rf   rg   rh   r   )rb   rc   rd   re   r3   rf   s         r
   adg_w_norm_forwardr      s(     w&(%","&"&( (r   c           	      &    t        | ||||dd      S )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)rg   rh   r   )rb   rc   rd   re   r3   s        r
   adg_wo_clip_forwardr      s     w1BE>fkx}~~r   )gX%?FT)gX%?)r   torch.nn.functionalr'   r(   Fr   r   r0   r   rF   rK   rQ   ra   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   