
    OiY!              
          d Z ddlZddlmZmZmZ ddl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 d	Z	 dde
dededefdZ	 dde
dedee   de
fdZde
dedededef
dZ	 	 	 ddedej2                  deeef   fdZy# e$ r d
ZY Tw xY w)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                 <   t        |      }t        j                  |d       t        | d      rkt        | j                  d      rUt        j
                  j                  |d      }| j                  j                  |       t        j                  d|        |S |r^t        j
                  j                  |d      }t        j                  | j                         |       t        j                  d|        |S i }| j                         D ]'  \  }}d	|v s|j                  j                         ||<   ) |st        j                   d
       yt        j
                  j                  |d      }t        j                  ||       t        j                  d|        |S )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	            W/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/training/lora_checkpoint.pysave_lora_weightsr+      sH    :&JKK
T*ui WU]]<M%Nww||J	:%%l3,\N;<	WW\\*j9


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

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

?I.,YK89    r)   _lora_configc                    t        |      }t        j                  j                  |      st	        d|       t        j                  j                  |      rPt        st        d      t        j                  | j                  |      | _
        t        j                  d|        | S |j                  d      rt        d      t        d|       )a  Load 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           	      R   t        |      }t        j                  |d       t        | |       |||j	                         |j	                         d}t        j
                  j                  |d      }t        j                  ||       t        j                  d| d| d| d       |S )	az  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*eZ( " ) 4 4 6 ) 4 4 6	N j*=>J	JJ~z*
KK
'
|8E7'+VWX r,   checkpoint_dirdevicec                 ,   dddddd}	 t        |       }t        j
                  j                  |d      }t        j
                  j                  |      r||d<   n$t        j
                  j                  |      r||d<   t        j
                  j                  |d      }t        j
                  j                  |      rB	 t        j                  ||d	
      }d|v r	 t        |d         |d<   d|v r	 t        |d         |d<   |d|v r	 |d   }
|l|
j                  di       j                         D ]I  }|j                         D ]4  \  }}t!        |t        j"                        s!|j%                  |      ||<   6 K |j'                  |
       d	|d<   t        j(                  d       |3d|v r/	 |j'                  |d          d	|d<   t        j(                  d       t        j(                  d|d    d|d           |S ddl}|j3                  d|      }|r9t        |j5                  d            |d<   t        j(                  d|d    d       |S # t        $ r t        j                  d|        |cY S w xY w# t        t        f$ r$}	t        j                  d|	 d       Y d}	~	d}	~	ww xY w# t        t        f$ r$}	t        j                  d|	 d       Y d}	~	d}	~	ww xY w# t*        t        t,        f$ r#}	t        j                  d|	        Y d}	~	rd}	~	ww xY w# t*        t        t,        f$ r#}	t        j                  d|	        Y d}	~	wd}	~	ww xY w# t.        t*        t        f$ r#}	t        j                  d|	        Y d}	~	|S d}	~	ww xY w) aO  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(-1#-a#>/0ttF|E!H )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   I ?M I< -M 2J2 M 
AK( %AK( (M /.L" !M #I98I9<J/J*$M *J//M 2K%K M  K%%M (L<LM LM "M6MM MM N0NN)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