
    ^j                         d Z ddlZddlZddlZddlmZ ddlZddlmZ  ej                  e
      Zd	dZ	 	 	 	 d
dee   fdZd	dZy)z+Checkpoint helpers for task-based training.    N)Optional)clean_state_dictc                 *   |xr t        t        j                  d      }|rPt        j                  j                  t        j
                  g      5  t        j                  | d|      cd d d        S t        j                  | d|      S # 1 sw Y   !xY w)Nsafe_globalscpu)map_locationweights_only)hasattrtorchserializationr   argparse	Namespaceload)checkpoint_pathr	   use_safe_globalss      ]/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/task/_helpers.py_load_train_checkpointr      s}    #T0C0C^(T  --x/A/A.BC 	^::oEP\]	^ 	^::oEUU	^ 	^s   B		Breturnc                    d}t         j                  j                  |      s.t        j	                  dj                  |             t               t        ||      }t        |t              rd}d}	d|v rd}d}	nd|v rd}|r|rt        j                  d       | j                  t        ||         |	r|j                  |	      nd       |/d	|v r+|rt        j                  d
       |j                  |d	          |C|j                  |v r5|rt        j                  d       |j                  ||j                            d|v r@|d   }d|v r|d   dkD  r|dz  }|r(t        j                  dj                  ||d                |S | j                  t        |             |r$t        j                  dj                  |             |S )zResume a task-based training checkpoint.

    Supports task checkpoints with ``state_dict``/``task_state`` and legacy
    training checkpoints that used a bare ``model`` key.
    NzNo checkpoint found at '{}'r	    
state_dict
task_statemodelz(Restoring model state from checkpoint...	optimizerz,Restoring optimizer state from checkpoint...z2Restoring AMP loss scaler state from checkpoint...epochversion   z!Loaded checkpoint '{}' (epoch {})zLoaded checkpoint '{}')ospathisfile_loggererrorformatFileNotFoundErrorr   
isinstancedictinfoload_checkpoint_stater   getload_state_dictstate_dict_key)
taskr   r   loss_scalerlog_infor	   resume_epoch
checkpointr,   task_state_keys
             r   resume_task_checkpointr3      s    L77>>/*3::?KL!!'lSJ*d#:%)N)N
"$NGH&& N!;<2@
~.d
 $
)BLL!OP))*[*AB&;+E+E+SLL!UV++J{7Q7Q,RS*$)'2
*z)/Dq/H A%LLL!D!K!KO]gho]p!qr/
;<-44_EF    c                    t        ||      }d}d}t        |t              r;|j                  dd      d}d}n$|j                  dd      d}nd|v rd}d}nd	|v rd	}t	        |r||   n|      }| j                  |t        |t              r|r|j                  |      ndd
       t        j                  dj                  |xs d|             y)z@Load EMA weights and optional task state into a task EMA module.r   r   state_dict_emaNtask_state_ema	model_emar   r   r   T)emazLoaded {} from checkpoint '{}'r1   )	r   r&   r'   r*   r   r)   r"   r(   r$   )r-   r   r	   r1   r,   r2   r   s          r   load_task_ema_checkpointr:   S   s    'lSJNN*d#>>*D1=-N-N^^K.:(NZ')N)N
"$N!*^"<T^_J*4Z*F>
~&_c  
 LL1889W<Yhijr4   )T)NNTT)__doc__r   loggingr   typingr   r   timm.modelsr   	getLogger__name__r"   r   intr3   r:    r4   r   <module>rC      s\    1   	   ( '

H
%V 9 c]9xkr4   