
    ^ji'                        d Z ddlZddlZddlZddlZddlmZmZmZm	Z	m
Z
 ddlZ	 ddlZdZ ej                   e      Zg dZdedefd	Z	 	 d ded
e
eej,                  f   defdZdededefdZdeeef   deeef   fdZ	 	 	 d!dedede
eej,                  f   dedeeef   f
dZ	 	 	 	 	 	 d"dej8                  j:                  dedede
eej,                  f   dedede	e   dedefdZ	 d#deeef   dej8                  j:                  dedeeef   fdZ	 	 	 	 d$dej8                  j:                  dede	ej@                  jB                     de	e   dedede	e"   fdZ#y# e$ r dZY >w xY w)%zi Model creation / weight loading / state_dict helpers

Hacked together by / Copyright 2020 Ross Wightman
    N)AnyCallableDictOptionalUnionTF)clean_state_dictload_checkpointload_state_dictremap_state_dictresume_checkpointcheckpoint_pathreturnc                     t        t        j                  d      sy	 t        j                  j                  t	        |             }|rddj                  |       dS dS # t
        $ r g }Y &w xY w)N get_unsafe_globals_in_checkpoint z Unsupported globals: z, .)hasattrtorchserializationr   str	Exceptionjoin)r   unsafe_globalss     _/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/_helpers.py_checkpoint_unsafe_globalsr      ss    5&&(JK,,MMcRaNbc ES#DIIn$=#>a@ZXZZ  s   (A A,+A,map_locationweights_onlyc                 *   |xr t        t        j                  d      }	 |rPt        j                  j                  t        j
                  g      5  t        j                  | ||      cd d d        S t        j                  | ||      S # 1 sw Y   !xY w# t        $ rA}|st        j                  | |      cY d }~S t        dt        j                   d      |d }~wt        j                  $ r!}|s t        dt        |        d      |d }~ww xY w)Nsafe_globalsr   r   )r   zAweights_only=True is not supported by this PyTorch build (torch==z). No automatic unsafe pickle fallback is performed. Upgrade PyTorch, or explicitly set weights_only=False only for trusted local checkpoints.zeweights_only=True blocked loading this checkpoint because it requires non-allowlisted pickle globals.zp No automatic unsafe pickle fallback is performed. If this checkpoint is trusted, retry with weights_only=False.)r   r   r   r   argparse	Namespaceload	TypeErrorRuntimeError__version__pickleUnpicklingErrorr   )r   r   r   use_safe_globalses        r   _torch_loadr+   )   s   
 $T0C0C^(T$$1183E3E2FG izz/[ghi izz/S_``i i  ::oLIIOPUPaPaOb ch h
 		
 !! s)/:; <LL

 	sM   0B B
(	B 2B 
BB 	DC7D=CD1DDtextprefixc                 D    | j                  |      r| t        |      d  S | S )N)
startswithlen)r,   r-   s     r   _remove_prefixr1   H   s$    vCKL!!K    
state_dictc                 n    i }d}| j                         D ]  \  }}|D ]  }t        ||      } |||<    |S )N)zmodule.z
_orig_mod.)itemsr1   )r3   cleaned_state_dict	to_removekvrs         r   r   r   O   s[    I   " "1 	%Aq!$A	% !1" r2   use_emadevicec                 Z   | rt         j                  j                  |       rt        |       j	                  d      r/t
        sJ d       t        j                  j                  | |      }nt        | ||      }d}t        |t              r;|r|j                  dd      d}n$|r|j                  dd      d}nd	|v rd	}nd
|v rd
}t        |r||   n|      }t        j                  dj!                  ||              |S t        j#                  dj!                  |              t%               )a_  Load state dictionary from checkpoint file.

    Args:
        checkpoint_path: Path to checkpoint file.
        use_ema: Whether to use EMA weights if available.
        device: Device to load checkpoint to.
        weights_only: Whether to load only weights (torch.load parameter).

    Returns:
        State dictionary loaded from checkpoint.
    z.safetensorsz-`pip install safetensors` to use .safetensors)r<   r    r   state_dict_emaN	model_emar3   modelzLoaded {} from checkpoint '{}'No checkpoint found at '{}')ospathisfiler   endswith_has_safetensorssafetensorsr   	load_filer+   
isinstancedictgetr   _loggerinfoformaterrorFileNotFoundError)r   r;   r<   r   
checkpointstate_dict_keyr3   s          r   r
   r
   ]   s   " 277>>/:((8#T%TT#$**44_V4TJ$_6XdeJj$':>>*:DAM!1Z^^K>J!,+!-J&!(%Nj&@Xbc
5<<^_]^3::?KL!!r2   r@   strictremap	filter_fnc                 4   t         j                  j                  |      d   j                         dv r)t	        | d      r| j                  |       yt        d      t        ||||      }|rt        ||       }n|r	 |||       }| j                  ||      }	|	S )a<  Load checkpoint into model.

    Args:
        model: Model to load checkpoint into.
        checkpoint_path: Path to checkpoint file.
        use_ema: Whether to use EMA weights if available.
        device: Device to load checkpoint to.
        strict: Whether to strictly enforce state_dict keys match.
        remap: Whether to remap state dict keys by order.
        filter_fn: Optional function to filter state dict.
        weights_only: Whether to load only weights (torch.load parameter).

    Returns:
        Incompatible keys from model.load_state_dict().
    )z.npzz.npyload_pretrainedz"Model cannot load numpy checkpointN)r<   r   )rS   )	rB   rC   splitextlowerr   rX   NotImplementedErrorr
   r   )
r@   r   r;   r<   rS   rT   rU   r   r3   incompatible_keyss
             r   r	   r	      s    2 
ww(,2248HH5+,!!/2 	 &&JKK !'&WcdJ%j%8
	z51
--j-Hr2   allow_reshapec                    i }t        |j                         j                         | j                               D ]  \  \  }}\  }}|j                         |j                         k(  s(J d| d|j                   d| d|j                   d	       |j                  |j                  k7  rF|r|j                  |j                        }n(J d| d|j                   d| d|j                   d	       |||<    |S )a  Remap checkpoint by iterating over state dicts in order (ignoring original keys).

    This assumes models (and originating state dict) were created with params registered in same order.

    Args:
        state_dict: State dict to remap.
        model: Model whose state dict keys to use.
        allow_reshape: Whether to allow reshaping tensors to match.

    Returns:
        Remapped state dictionary.
    zTensor size mismatch z: z vs z. Remap failed.zTensor shape mismatch )zipr3   r5   numelshapereshape)r3   r@   r]   out_dictkavakbvbs           r   r   r      s    " H!%"2"2"4":":"<j>N>N>PQ R(2rxxzRXXZ't+@BrxxjPTUWTXXZ[][c[cZdds)tt'88rxxZZ)j 6rd"RXXJd2$bQSQYQYPZZijju Or2   	optimizerloss_scalerlog_infoc                 <   d}t         j                  j                  |      rMt        |d|      }t	        |t
              rd|v r|rt        j                  d       t        |d         }| j                  |       |/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                  |       |r$t        j                  dj                  |             |S t        j                  dj                  |             t               )a  Resume training from checkpoint.

    Args:
        model: Model to load checkpoint into.
        checkpoint_path: Path to checkpoint file.
        optimizer: Optional optimizer to restore state.
        loss_scaler: Optional AMP loss scaler to restore state.
        log_info: Whether to log loading info.
        weights_only: Whether to load only weights via torch.load.

    Returns:
        Resume epoch number if available, else None.
    Ncpur    r3   z(Restoring model state from checkpoint...rh   z,Restoring optimizer state from checkpoint...z2Restoring AMP loss scaler state from checkpoint...epochversion   z!Loaded checkpoint '{}' (epoch {})zLoaded checkpoint '{}'rA   )rB   rC   rD   r+   rI   rJ   rL   rM   r   r
   rR   rN   rO   rP   )	r@   r   rh   ri   rj   r   resume_epochrQ   r3   s	            r   r   r      s   * L	ww~~o& uS_`
j$'LJ,FGH)*\*BCJ!!*-$
)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
  !!*-5<<_MN3::?KL!!r2   )rl   T)Trl   T)Trl   TFNT)T)NNTT)$__doc__r!   loggingrB   r'   typingr   r   r   r   r   r   safetensors.torchrG   rF   ImportError	getLogger__name__rL   __all__r   r   r<   boolr+   r1   r   r
   nnModuler	   r   optim	Optimizerintr    r2   r   <module>r      s     	  7 7  '

H
%[ [ [ 27!C-. > c c c3h DcN   +0!	("("(" c5<<'((" 	("
 
#s(^("\ +0(,!'xx'' ' c5<<'(	'
 ' ' H%' ' 	'Z #cNxx  
#s(^	@ 6:%)!6"xx6"6" EKK1126" c]	6"
 6" 6" c]6"}  s   E   E+*E+