
    xi!                        d Z ddlZddlZddlmZmZmZmZ ddlZddl	m
Z
 ddlmZ ddlmZ 	 ddlmZmZ dZdefdZdedefdZ	 ddededeedeeef   f   fdZ	 	 ddddedeej6                     deeeef      def
dZdddedeeef   fdZ	 	 ddddedededee   deeeef      defdZy# e$ r d	ZeZ e
j(                  d
       Y w xY w) 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                      t         S )zCheck if LyCORIS is importable.)LYCORIS_AVAILABLE     :/mnt/workspace/ACE-Step-1.5/acestep/training/lokr_utils.pycheck_lycoris_availabler      s    r   module_namec                     | sy| j                         }|xs g D ]M  }t        |      j                         j                         }|s-|j                  |      sd| |v sd| |v sM y y)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   #   su    D &B&K%%'==1#w$AaS'T/ ' r   lokr_config
multiplierr
   c                 
   t         st        d      | j                  }| j                         D ]  \  }}d|_         t        j                  |j                  |j                  d       t        |||j                  |j                  d|j                  |j                  |j                  |j                  |j                  |j                   |j"                  |j$                        }|j&                  r~	 t        |||j                  |j                  d|j                  |j                  |j                  |j                  |j                  |j                   |j"                  |j$                  d      }|j/                          ||_        g }d
}	d
}
g }t3        t5        |dg       xs g       D ]  \  }}t5        |dd	      xs* t5        |dd	      xs |j6                  j8                   d| }t;        ||j                        }|r|	dz  }	n$|
dz  }
t=        |      dk  r|j?                  |       |jA                         D ]  }||_        |s|j?                  |         t+        jB                  d|	 d|
 d|j                          |r't+        jB                  ddjE                  |      z          |s-|jA                         D ]  }d|_        |j?                  |        |D ci c]  }tG        |      | }}tI        d | jA                         D              }tI        d |jK                         D              }tI        d |jK                         D              }||||d
kD  r||z  nd|j                  |j                  |j                  d|j                  d	}t+        jB                  d       t+        jB                  d|dd|dd|d    d!d"       | ||fS # t(        $ r#}t+        j,                  d|        Y d	}~d	}~ww xY wc c}w )#zq
    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   <   K   | ]  }|j                           y wNnumel.0ps     r   	<genexpr>z'inject_lokr_into_dit.<locals>.<genexpr>   s     =*<Qqwwy*<   c              3   <   K   | ]  }|j                           y wr8   r9   r;   s     r   r>   z'inject_lokr_into_dit.<locals>.<genexpr>   s     @)?Aaggi)?r?   c              3   V   K   | ]!  }|j                   s|j                          # y wr8   )requires_gradr:   r;   s     r   r>   z'inject_lokr_into_dit.<locals>.<genexpr>   s     X.D1779.Ds   ))g        )	total_paramslokr_paramstrainable_paramstrainable_ratior&   r'   r)   r(   r   zLoKr injected into decoderzLoKr trainable params: ,/z (rF   z.2%))&r   ImportErrordecodernamed_parametersrB   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!   rK   r   paramlycoris_netexclokr_param_listenabled_module_countdisabled_module_countdisabled_examplesidxmoduler   enabledr=   unique_paramsrC   rD   rE   rZ   s                         r   inject_lokr_into_ditrk   1   s'    5
 	

 mmG **,5# -  + : :&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&  'GO gr!B!HbIVFK. 4vvt,4""++,AcU3 	
 .k;;U;UV A% !Q&!$%)!((5&&(E")E&&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   2A=O P 	O=O88O=ra   
output_dirdtypemetadatac                    t        |      }t        j                  |d       t        j                  j	                  |d      }ddd}|rK|j                         D ]8  \  }}|	t        |t              r|||<   t        j                  |d      ||<   : | j                  |||       t        j                  d	|        |S )
z!Save LoKr weights to safetensors.Texist_okzlokr_weights.safetensorsr%   lycoris)r(   format)ensure_ascii)rm   rn   zLoKr weights saved to )r	   osmakedirspathr[   items
isinstancer   jsondumpssave_weightsr   rZ   )ra   rl   rm   rn   weights_pathsave_metadatakeyvalues           r   save_lokr_weightsr      s     :&JKK
T*77<<
,FGL-3y$IM"..*JC}%%%*c"%)ZZD%Ic" + \O
KK(78r   r}   c                     t        |      }t        j                  j                  |      st	        d|       | j                  |      }t        j                  d|        |S )z3Load LoKr weights into an injected LyCORIS network.zLoKr weights not found: zLoKr weights loaded from )r	   ru   rw   existsFileNotFoundErrorload_weightsr   rZ   )ra   r}   results      r   load_lokr_weightsr      sX    \*L77>>,'":<. IJJ%%l3F
KK+L>:;Mr   epochglobal_steprun_metadatac           	         t        |      }t        j                  |d       i }||j                         |d<   |||d<   |xs d}t	        | ||       |||j                         |j                         d}	||j                         |	d<   |||	d<   t        j                  j                  |d      }
t        j                  |	|
       t        j                  d	| d
| d| d       |S )z1Save LoKr weights plus optimizer/scheduler state.Trp   Nr    r   )rn   )r   r   optimizer_state_dictscheduler_state_dictztraining_state.ptzLoKr checkpoint saved to z (epoch=z, step=rI   )r	   ru   rv   to_dictr   
state_dictrw   r[   torchsaver   rZ   )ra   	optimizer	schedulerr   r   rl   r    r   rn   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	JJuj!
KK+J<xwgk]Z[\]r   )g      ?)NN) __doc__rz   ru   typingr   r   r   r   r   logurur   acestep.training.configsr   acestep.training.path_safetyr	   rr   r
   r   r   rJ   rP   boolr   r   r   floatrk   rm   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FNN	1s   
B3 3CC