
    xi!                        S r SSKrSSKrSSKJrJrJrJr  SSKrSSK	J
r
  SSKJr  SSKJr   SSKJrJr  SrS\4S jrS\S\4S jr SS\S\S\\S\\\4   4   4S jjr  SSSS\S\\R6                     S\\\\4      S\4
S jjrSSS\S\\\4   4S jr  SSSS\S\S\S\\   S\\\\4      S\4S jjrg! \ a    S	r\r\
R(                  " S
5         Nf = f) z~
LoKr utilities for ACE-Step training and inference.

This module integrates LyCORIS LoKr adapters with the ACE-Step decoder.
    N)AnyDictOptionalTuple)logger)
LoKRConfig)	safe_path)LycorisNetworkcreate_lycorisTFzjLyCORIS library not installed. LoKr training/inference unavailable. Install with: pip install lycoris-lorareturnc                      [         $ )zCheck if LyCORIS is importable.)LYCORIS_AVAILABLE     R/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/training/lokr_utils.pycheck_lycoris_availabler      s    r   module_namec                    U (       d  gU R                  5       nU=(       d    /  H]  n[        U5      R                  5       R                  5       nU(       d  M3  UR                  U5      (       d  SU 3U;   d  SU 3U;   d  M]    g   g)zKReturn True if a LyCORIS module name maps to one of target module suffixes.F_.T)lowerstrstripendswith)r   target_modulesnametargetts        r   _matches_target_module_namer   #   sx    D &B&K%%'==1#w$AaS'T/ ' r   lokr_config
multiplierr
   c                 v   [         (       d  [        S5      eU R                  nU R                  5        H  u  pESUl        M     [
        R                  " UR                  UR                  S.5        [        UUUR                  UR                  SUR                  UR                  UR                  UR                  UR                  UR                   UR"                  UR$                  S9nUR&                  (       a{   [        UUUR                  UR                  SUR                  UR                  UR                  UR                  UR                  UR                   UR"                  UR$                  SS9nUR/                  5         Xcl        / nS
n	S
n
/ n[3        [5        US/ 5      =(       d    / 5       H  u  p[5        USS	5      =(       d/    [5        USS	5      =(       d    UR6                  R8                   SU 3n[;        XR                  5      nU(       a  U	S-  n	O%U
S-  n
[=        U5      S:  a  UR?                  U5        URA                  5        H#  nXl        U(       d  M  UR?                  U5        M%     M     [*        RB                  " SU	 SU
 SUR                   35        U(       a(  [*        RB                  " SSRE                  U5      -   5        U(       d/  URA                  5        H  nSUl        UR?                  U5        M     U Vs0 s H  n[G        U5      U_M     nn[I        S U RA                  5        5       5      n[I        S URK                  5        5       5      n[I        S URK                  5        5       5      nUUUUS
:  a  UU-  OSUR                  UR                  UR                  SUR                  S.	n[*        RB                  " S5        [*        RB                  " SUS SUS SUS    S! S"35        XU4$ ! [(         a$  n[*        R,                  " SU 35         S	nAGNS	nAff = fs  snf )#za
Inject LoKr adapters into the decoder.

Returns:
    Tuple: (model, lycoris_network, info_dict)
zULyCORIS library is required for LoKr training. Install with: pip install lycoris-loraF)unet_target_nametarget_namelokr)
linear_dimlinear_alphaalgofactordecompose_both
use_tucker
use_scalarfull_matrixbypass_moders_loraunbalanced_factorizationT)r&   r'   r(   r)   r*   r+   r,   r-   r.   r/   r0   dora_wdz2DoRA mode not supported in current LyCORIS build: Nr   loras	lora_namer   #      zLoKr target filter: enabled z LyCORIS modules (disabled z) for targets=z+LoKr disabled non-target modules (sample): z, c              3   @   #    U  H  oR                  5       v   M     g 7fNnumel.0ps     r   	<genexpr>'inject_lokr_into_dit.<locals>.<genexpr>   s     =*<Qwwyy*<   c              3   @   #    U  H  oR                  5       v   M     g 7fr8   r9   r;   s     r   r>   r?      s     @)?Aggii)?r@   c              3   f   #    U  H'  oR                   (       d  M  UR                  5       v   M)     g 7fr8   )requires_gradr:   r;   s     r   r>   r?      s      X.D917799.Ds   11g        )	total_paramslokr_paramstrainable_paramstrainable_ratior&   r'   r)   r(   r   zLoKr injected into decoderzLoKr trainable params: ,/z (rG   z.2%))&r   ImportErrordecodernamed_parametersrC   r
   apply_presetr   r   r&   r'   r)   r*   r+   r,   r-   r.   r/   r0   weight_decompose	Exceptionr   warningapply_to_lycoris_net	enumerategetattr	__class____name__r   lenappend
parametersinfojoinidsumvalues)modelr    r!   rL   r   paramlycoris_netexclokr_param_listenabled_module_countdisabled_module_countdisabled_examplesidxmoduler   enabledr=   unique_paramsrD   rE   rF   r[   s                         r   inject_lokr_into_ditrl   1   s     5
 	

 mmG **,# -  + : :&55	
 !)) --!!"11))))++++##!,!E!EK  ##	W(&11(55"))*99&11&11'33'33#++)4)M)MK&  'O gr!B!HbIFK. 4vvt,4""++,AcU3 	
 .k;U;UV A% !Q&!$%)!((5&&(E")w&&u- ) J( KK
&';&< =*+>+:T:T9U	W ADIIN_D``a ++-E"&E""5) .
 (77!RUAXM7=%*:*:*<==L@)=)=)?@@KXm.B.B.DXX %",>JQ>N+l:TW!,,#00$$%44
D KK,-
KK
!"21!5Q|A6F G"#C(	+ t##K  	WNNOPSuUVV	W^ 8s   ;A:P P6
P3P..P3rb   
output_dirdtypemetadatac                    [        U5      n[        R                  " USS9  [        R                  R	                  US5      nSSS.nU(       aP  UR                  5        H<  u  pgUc  M
  [        U[        5      (       a  XuU'   M%  [        R                  " USS9XV'   M>     U R                  XBUS9  [        R                  " S	U 35        U$ )
z!Save LoKr weights to safetensors.Texist_okzlokr_weights.safetensorsr%   lycoris)r(   format)ensure_ascii)rn   ro   zLoKr weights saved to )r	   osmakedirspathr\   items
isinstancer   jsondumpssave_weightsr   r[   )rb   rm   rn   ro   weights_pathsave_metadatakeyvalues           r   save_lokr_weightsr      s     :&JKK
T*77<<
,FGL-3y$IM"..*JC}%%%%*c"%)ZZD%I" + \O
KK(78r   r~   c                     [        U5      n[        R                  R                  U5      (       d  [	        SU 35      eU R                  U5      n[        R                  " SU 35        U$ )z3Load LoKr weights into an injected LyCORIS network.zLoKr weights not found: zLoKr weights loaded from )r	   rv   rx   existsFileNotFoundErrorload_weightsr   r[   )rb   r~   results      r   load_lokr_weightsr      s[    \*L77>>,''":<. IJJ%%l3F
KK+L>:;Mr   epochglobal_steprun_metadatac           	         [        U5      n[        R                  " USS9  0 nUb  UR                  5       US'   Ub  XxS'   U=(       d    Sn[	        XUS9  UUUR                  5       UR                  5       S.n	Ub  UR                  5       U	S'   Ub  XyS'   [        R                  R                  US5      n
[        R                  " X5        [        R                  " S	U S
U SU S35        U$ )z1Save LoKr weights plus optimizer/scheduler state.Trq   Nr    r   )ro   )r   r   optimizer_state_dictscheduler_state_dictztraining_state.ptzLoKr checkpoint saved to z (epoch=z, step=rJ   )r	   rv   rw   to_dictr   
state_dictrx   r\   torchsaver   r[   )rb   	optimizer	schedulerr   r   rm   r    r   ro   state
state_paths              r   save_lokr_training_checkpointr      s     :&JKK
T*!H"-"5"5"7#/ 4HkA " ) 4 4 6 ) 4 4 6	E *224m ,nj*=>J	JJu!
KK+J<xwgk]Z[\]r   )g      ?)NN) __doc__r{   rv   typingr   r   r   r   r   logurur   acestep.training.configsr   acestep.training.path_safetyr	   rs   r
   r   r   rK   rQ   boolr   r   r   floatrl   rn   r   r   intr   r   r   r   <module>r      s    	 - -   / 2
6 
S T " E$E$ E$ 3 $sCx.01	E$V $()-	! EKK  tCH~&	
 	4#3 3 4PSUXPX> " )--1$!$ 	$
 $ $ *%$ 4S>*$ 	$Q  N
NN	1s   
B< <CC