
    OiY!              
       ,   S r SSKrSSK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SKJr   SSKJr  S	r SS\
S\S\S\4S jjr SS\
S\S\\   S\
4S jjrS\
S\S\S\S\4
S jr   SS\S\R2                  S\\\4   4S jjrg! \ a    S
r N\f = f)ze
LoRA Checkpoint Utilities for ACE-Step

Provides functions for saving and loading LoRA checkpoints.
    N)OptionalDictAny)logger)Module)	safe_path)
LoRAConfig)	PeftModelTFmodel
output_dirsave_full_modelreturnc                 p   [        U5      n[        R                  " USS9  [        U S5      (       aq  [        U R                  S5      (       aV  [        R
                  R                  US5      nU R                  R                  U5        [        R                  " SU 35        U$ U(       a`  [        R
                  R                  US5      n[        R                  " U R                  5       U5        [        R                  " SU 35        U$ 0 nU R                  5        H)  u  pgS	U;   d  M  UR                  R                  5       XV'   M+     U(       d  [        R                   " S
5        g[        R
                  R                  US5      n[        R                  " XX5        [        R                  " SU 35        U$ )zSave LoRA adapter weights.

Args:
    model: Model with LoRA adapters
    output_dir: Directory to save weights
    save_full_model: Whether to save the full model state dict

Returns:
    Path to saved weights
Texist_okdecodersave_pretrainedadapterzLoRA adapter saved to zmodel.ptzFull model state dict saved to lora_z!No LoRA parameters found to save! zlora_weights.ptzLoRA weights saved to )r   osmakedirshasattrr   pathjoinr   r   infotorchsave
state_dictnamed_parametersdataclonewarning)	r   r   r   adapter_path
model_pathlora_state_dictnameparam	lora_paths	            ?/mnt/workspace/ACE-Step-1.5/acestep/training/lora_checkpoint.pysave_lora_weightsr+      sH    :&JKK
T*ui  WU]]<M%N%Nww||J	:%%l3,\N;<	WW\\*j9


5##%z25j\BC 113KD$(-

(8(8(:% 4 NN>?GGLL->?	

?.,YK89    r)   _lora_configc                    [        U5      n[        R                  R                  U5      (       d  [	        SU 35      e[        R                  R                  U5      (       aW  [        (       d  [        S5      e[        R                  " U R                  U5      U l
        [        R                  " SU 35        U $ UR                  S5      (       a  [        S5      e[        SU 35      e)zLoad LoRA adapter weights into the model.

Args:
    model: The base model (without LoRA)
    lora_path: Path to saved LoRA adapter directory
    lora_config: Unused; retained for API compatibility

Returns:
    Model with LoRA weights loaded
zLoRA weights not found: zHPEFT library is required to load adapter. Install with: pip install peftzLoRA adapter loaded from z.ptzcLoading LoRA weights from .pt files is disabled for security. Use a PEFT adapter directory instead.z Unsupported LoRA weight format: )r   r   r   existsFileNotFoundErrorisdirPEFT_AVAILABLEImportErrorr
   from_pretrainedr   r   r   endswith
ValueError)r   r)   r-   	validateds       r*   load_lora_weightsr8   E   s     )$I77>>)$$":9+ FGG	ww}}Y~Z  "11%--K/	{;< L 
		E	"	"4
 	
 ;I;GHHr,   epochglobal_stepc           	      N   [        U5      n[        R                  " USS9  [        X5        UUUR	                  5       UR	                  5       S.n[        R
                  R                  US5      n[        R                  " Xg5        [        R                  " SU SU SU S35        U$ )	aR  Save a training checkpoint including LoRA weights and training state.

Args:
    model: Model with LoRA adapters
    optimizer: Optimizer state
    scheduler: Scheduler state
    epoch: Current epoch number
    global_step: Current global step
    output_dir: Directory to save checkpoint

Returns:
    Path to saved checkpoint directory
Tr   )r9   r:   optimizer_state_dictscheduler_state_dicttraining_state.ptzTraining checkpoint saved to z (epoch , step ))r   r   r   r+   r   r   r   r   r   r   r   )r   	optimizer	schedulerr9   r:   r   training_state
state_paths           r*   save_training_checkpointrE   m   s    * :&JKK
T*e( " ) 4 4 6 ) 4 4 6	N j*=>J	JJ~*
KK
'
|8E7'+VWX r,   checkpoint_dirdevicec                    SSSSSS.n [        U 5      n[        R
                  R                  US5      n[        R
                  R                  U5      (       a  XdS'   O([        R
                  R                  U5      (       a  XTS'   [        R
                  R                  US5      n[        R
                  R                  U5      (       GaU   [        R                  " XsS	S
9nSU;   a   [        US   5      US'   SU;   a   [        US   5      US'   Ub  SU;   a   US   n
Ubt  U
R                  S0 5      R                  5        HP  nUR                  5        H9  u  p[!        U[        R"                  5      (       d  M&  UR%                  U5      X'   M;     MR     UR'                  U
5        S	US'   [        R(                  " S5        Ub6  SU;   a0   UR'                  US   5        S	US'   [        R(                  " S5        [        R(                  " SUS    SUS    35        U$ SSKnUR3                  SU5      nU(       a:  [        UR5                  S5      5      US'   [        R(                  " SUS    S35        U$ ! [         a    [        R                  " SU < 35        Us $ f = f! [        [        4 a%  n	[        R                  " SU	 S35         Sn	A	GNSn	A	ff = f! [        [        4 a%  n	[        R                  " SU	 S35         Sn	A	GNSn	A	ff = f! [*        [        [,        4 a$  n	[        R                  " SU	 35         Sn	A	GNSn	A	ff = f! [*        [        [,        4 a$  n	[        R                  " SU	 35         Sn	A	GNSn	A	ff = f! [.        [*        [        4 a$  n	[        R                  " SU	 35         Sn	A	U$ Sn	A	ff = f) a  Load training checkpoint.

Args:
    checkpoint_dir: Directory containing checkpoint files
    optimizer: Optimizer instance to load state into (optional).
        When provided, loads optimizer_state_dict from the checkpoint.
    scheduler: Scheduler instance to load state into (optional).
        When provided, loads scheduler_state_dict from the checkpoint.
    device: Device to load tensors to

Returns:
    Dictionary with checkpoint info:
    - epoch: Saved epoch number
    - global_step: Saved global step
    - adapter_path: Path to adapter weights
    - loaded_optimizer: Whether optimizer state was loaded (True when optimizer param provided and state loaded)
    - loaded_scheduler: Whether scheduler state was loaded (True when scheduler param provided and state loaded)
r   NF)r9   r:   r$   loaded_optimizerloaded_schedulerz&Rejected unsafe checkpoint directory: r   r$   r>   T)map_locationweights_onlyr9   z0Failed to parse 'epoch' from training_state.pt: z, using default 0r:   z6Failed to parse 'global_step' from training_state.pt: r<   staterI   z&Loaded optimizer state from checkpointz Failed to load optimizer state: r=   rJ   z&Loaded scheduler state from checkpointz Failed to load scheduler state: z&Loaded checkpoint metadata from epoch r?   z"Failed to load training_state.pt: zepoch_(\d+)   z,No training_state.pt found, extracted epoch z
 from path)r   r6   r   r#   r   r   r   r1   isfiler   loadint	TypeErrorgetvaluesitems
isinstanceTensortoload_state_dictr   RuntimeErrorKeyErrorOSErrorresearchgroup)rF   rA   rB   rG   resultsafe_dirr$   rD   rC   eoptimizer_staterM   kvr]   matchs                   r*   load_training_checkpointrg      s   2 !!F^,
 77<<)4L	ww}}\""!-~	x	 	 !)~h(;<J	ww~~j!!.	E"ZZdN .(&).*A&BF7O
 .,/}0M,NF=) $)?>)QK&45K&LO)%4%8%8"%E%L%L%NE(-#-a#>#>/0ttF|EH )6 &O --o>15F-.KK HI $)?>)QK--n=S.TU15F-.KK HI KK88IQWXeQfPgh M 			.(3!%++a.1F7OKK>vg>OzZ MK  ??QRS( #I. NNJ1#M^_  #I. NNPQRPSSde   %j(; KNN%EaS#IJJK %j(; KNN%EaS#IJJK z2 	ENN?sCDD M	Es   I9 N (J% 9N  K 	N AL <AL 	N /M <"N 9&J"!J"%K5KN KN L-LN LN M*M	N 	MN N
&N?N N

N O"OO)F)N)NNN)__doc__r   typingr   r   r   logurur   r   torch.nnr   acestep.training.path_safetyr   acestep.training.configsr	   peftr
   r2   r3   strboolr+   r8   rQ   rE   rG   rg    r,   r*   <module>rr      s   
 & &    2 /N "))) ) 		)^ *.%%% :&% 	%P'' 	'
 ' ' 	'X 	gg LL	g
 
#s(^gE  Ns   B BB