
    ^ie                     :   S r SSKJrJrJrJr  SSKJr  SSKrSSK	J
r
  SSKJr  S r SSKJrJrJrJr  SrS\
R,                  S\
R,                  4S jrS\\   4S jrSS\SS4S jjrS\S\\\\\4   4   4S jrg! \ a    S	r\R*                  " S
5         Nbf = f)zs
LoRA Injection Utilities for ACE-Step

Provides functions for injecting LoRA adapters into the DiT decoder model.
    )ListTupleAnyDict)loggerN)
LoRAConfigc                    [        U SS5      n Ub  U" 5       nOSn SU l        U$ ! [         a    [        R                  " SSS9   U$ f = f! [
         a     SU l        O%! [         a    [        R                  " SSS9   Of = f[        U SS5      (       dD  [        R                  " S5         SU l        O%! [         a    [        R                  " S	SS9   Of = f gf = f)
a%  Safely call enable_input_require_grads on the decoder.

This helper wraps the original enable_input_require_grads method,
handling NotImplementedError gracefully and tracking whether the hook
was successfully enabled.

Args:
    self: The decoder module to call enable_input_require_grads on.
(_acestep_orig_enable_input_require_gradsNTz/Failed to set _acestep_input_grads_hook_enabled)exc_infoF$_acestep_input_grads_warning_emittedzkSkipping enable_input_require_grads for decoder: get_input_embeddings is not implemented (expected for DiT)z2Failed to set _acestep_input_grads_warning_emitted)getattr!_acestep_input_grads_hook_enabled	Exceptionr   debugNotImplementedErrorinfor   )selforig_enable_input_require_gradsresults      V/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/training/lora_injection.py _safe_enable_input_require_gradsr      s     '.8$'#*646FF	59D2
 	  	LLAD 		
  	5:D2 	LLAD	 tCUKKKKM<@9 HSW %sy   A & A A	A A		A 
CACB>C B+C-B54C5CCCCC)get_peft_model
LoraConfigTaskType	PeftModelTFz@PEFT library not installed. LoRA training will not be available.modulereturnc                 |   U n[        U5      1n [        USS5      nUc  O&[        U5      nXB;   a  OUR                  U5        UnM7  [        USS5      nUb4  [        USS5      nUb"  [        U[        R
                  5      (       a  UnOUn[        USS5      nUb!  [        U[        R
                  5      (       a  UnU$ )aZ  Unwrap PEFT/Fabric wrappers from a model/decoder to retrieve the base DiT module.

This internal helper walks the wrapper chain and returns the underlying
``nn.Module`` that can be passed to PEFT for adapter injection.

Args:
    module: A model or decoder that may have PEFT/Fabric wrappers.

Returns:
    The unwrapped base DiT decoder module.
_forward_moduleN
base_modelmodel)idr   add
isinstancennModule)r   decoderseen_idsnext_decodernext_idr    inner_modelfinal_models           r   _unwrap_decoderr-   M   s     G7}H
w(94@\"W  ,5Jj'48"z+ryy'I'I!G G'7D1K:k299#E#EN    c                   ^ / n[        U S5      (       ar  U R                  R                  5        HT  u  mn[        U4S jS 5       5      (       d  M"  [	        U[
        R                  5      (       d  MC  UR                  T5        MV     U$ )zGet the list of module names in the DiT decoder that can have LoRA applied.

Args:
    model: The AceStepConditionGenerationModel

Returns:
    List of module names suitable for LoRA
r'   c              3   ,   >#    U  H	  oT;   v   M     g 7fN ).0projnames     r   	<genexpr>)get_dit_target_modules.<locals>.<genexpr>   s     U,TD4<,Ts   )q_projk_projv_projo_proj)hasattrr'   named_modulesanyr$   r%   Linearappend)r!   target_modulesr   r5   s      @r   get_dit_target_modulesrB   t   sk     Nui  !MM779LD&U,TUUUfbii00"))$/ :
 r.   freeze_encoderc                   ^	 SnU R                  5        H_  u  m	nST	;   nT	R                  U5      =(       d    [        U	4S jU 5       5      nU(       a	  SUl        MH  U(       d	  U(       a  MX  SUl        Ma     SnSnU R                  5        H<  u  pXcR	                  5       -  nUR                  (       d  M*  XsR	                  5       -  nM>     [
        R                  " SXg-
  S 35        [
        R                  " S	US 35        g
)zFreeze all non-LoRA parameters in the model.

Args:
    model: The model to freeze parameters for
    freeze_encoder: Whether to freeze the encoder (condition encoder)
)encodertext_encodervision_encoderzmodel.encoderlora_c              3   L   >#    U  H  nTR                  U S 35      v   M     g7f).N)
startswith)r3   prefixr5   s     r   r6   -freeze_non_lora_parameters.<locals>.<genexpr>   s'      >
8HfDOOvhaL))8Hs   !$TFr   zFrozen parameters: ,zTrainable parameters: N)named_parametersrK   r>   requires_gradnumelr   r   )
r!   rC   encoder_prefixesparamis_lora
is_encodertotal_paramstrainable_params_r5   s
            @r   freeze_non_lora_parametersrY      s     V--/eT/__%56 
# >
8H>
 ;

 "&E::"'E 0 L**,%- -
 KK%l&Ea%HIJ
KK()9!(<=>r.   lora_configc           	      Z   [         (       d  [        S5      e[        U R                  5      nX l        [	        US5      (       aC  [	        US5      (       d2  UR
                  nX2l        [        R                  " [        U5      Ul        [	        US5      (       a   SUl
        [        UR                  UR                  UR                  UR                   UR"                  [$        R&                  S9n[)        X$5      nXPl        U R+                  5        H  u  pgSU;  d  M  SUl        M     [/        S U R1                  5        5       5      n[/        S	 U R1                  5        5       5      n	UU	US
:  a  X-  OS
UR                  UR                  UR                   S.n
[2        R4                  " S5        [2        R4                  " SUS 35        [2        R4                  " SU	S SU
S   S S35        [2        R4                  " SUR                   SUR                   35        X
4$ ! [         a     GNf = f)zInject LoRA adapters into the DiT decoder of the model.

Args:
    model: The AceStepConditionGenerationModel
    lora_config: LoRA configuration

Returns:
    Tuple of (peft_model, info_dict)
zJPEFT library is required for LoRA training. Install with: pip install peftenable_input_require_gradsr
   is_gradient_checkpointingF)r
lora_alphalora_dropoutrA   bias	task_typerH   c              3   @   #    U  H  oR                  5       v   M     g 7fr1   )rQ   r3   ps     r   r6   'inject_lora_into_dit.<locals>.<genexpr>   s     =*<Qwwyy*<s   c              3   f   #    U  H'  oR                   (       d  M  UR                  5       v   M)     g 7fr1   )rP   rQ   rd   s     r   r6   rf      s      T.@OO917799.@s   11r   )rV   rW   trainable_ratiolora_rr_   rA   zLoRA injected into DiT decoder:z  Total parameters: rN   z  Trainable parameters: z (rh   z.2%)z  LoRA rank: z	, alpha: )PEFT_AVAILABLEImportErrorr-   r'   r<   r\   r
   types
MethodTyper   r]   r   r   r^   alphadropoutrA   ra   r   FEATURE_EXTRACTIONr   rO   rP   sum
parametersr   r   )r!   rZ   r'   origpeft_lora_configpeft_decoderr5   rS   rV   rW   r   s              r   inject_lora_into_ditrw      s    >X
 	
 emm,GMw455g;? ? 11;?8-2-=-=,g.
* w344	05G- "
--$$ (("11-- "'<L M--/$"'E 0 =%*:*:*<==LTe.>.>.@TT %,>JQ>N+:TU--!''%44D KK12
KK&|A&678
KK
"#3A"6b>O9PQT8UUVW KK-i8I8I7JKL;M  		s   H 
H*)H*)T)__doc__typingr   r   r   r   logurur   rm   torch.nnr%   acestep.training.configsr   r   peftr   r   r   r   rk   rl   warningr&   r-   strrB   boolrY   rw   r2   r.   r   <module>r      s    * )    /,^W  N$BII $")) $NT#Y (?d ?d ?@GG 3S#XGA  WN
NNUVWs   A= =BB