
      iO2                     8   d dl Z d dlZd dlZd dlmZ d dlmZ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mZ d dlmZ d dlmZ d d	lmZ d d
lmZ d dlmZ d dlmZ d dlm Z m!Z!  e jD                  e#      Z$dddedededef
dZ%ddde&ddfdZ'd)dZ(d)dZ)d)dZ*ddddde+deded   dedefdZ,ddde+dededef
dZ-dd ddde+ded!ee.   deddfd"Z/ddde0e+e0f   fd#Z1ddd$e0e+ef   ddfd%Z2ddd$e0e+ef   ddfd&Z3ddd$e0e+ef   ddfd'Z4ddde+dededef
d(Z5y)*    N)deepcopy)AnyCallableOptionalUnion)Version)_DeviceDtypeModuleMixin)
CheckpointEarlyStopping)WandbLogger)_SubprocessScriptLauncher)_get_sigkill_signal)TrainerStatus)_TunerExitException)is_overridden)rank_zero_inforank_zero_warntrainer
pl.Trainer
trainer_fnargskwargsreturnc                 <   	 | j                   j                  , | j                   j                  j                  |g|d| i|S  ||i |S # t        $ rN t	        |        | j                          t        j                  | j                  _	        d| j                  _
        Y yt        $ r}t        d       t        j                  t        j                  t        j                         t!        | |       | j                          | j                   j                  }t#        |t$              r|j'                  t)                      t+        j,                  d       Y d}~yd}~wt.        $ r3}t!        | |       | j                          d| j                  _
         d}~ww xY w)am  Error handling, intended to be used only for main trainer function entry points (fit, validate, test, predict)
    as all errors should funnel through them.

    Args:
        trainer_fn: one of (fit, validate, test, predict)
        *args: positional arguments to be passed to the `trainer_fn`
        **kwargs: keyword arguments to be passed to `trainer_fn`

    Nr   z=
Detected KeyboardInterrupt, attempting graceful shutdown ...   )strategylauncherlaunchr   _call_teardown_hook	_teardownr   FINISHEDstatestatusstageKeyboardInterruptr   signalSIGINTSIG_IGN
_interrupt
isinstancer   killr   sysexitBaseException)r   r   r   r   	exceptionr   s         s/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pytorch_lightning/trainer/call.py_call_and_handle_interruptr1   $   s>   $$037##,,33JaawaZ`aa4*6** #G$,55" 	WXfmmV^^47I&##,,h 9:MM-/0 7I&"s2   AA A AF"F*B-EF(.FFr/   c                    t         j                  | j                  _        t	        | d|       | j
                  t        | d|       | j                  j                  |       | j                  D ]  }|j                  d        y )Non_exceptionfailed)r   INTERRUPTEDr"   r#   _call_callback_hooks
datamodule_call_lightning_datamodule_hookr   r3   loggersfinalize)r   r/   loggers      r0   r)   r)   L   sf    (44GMM.)<%'K!!),//! "    c                 >   | j                   j                  J | j                   j                  }| j                  j                         D ].  }t	        |t
              s| j                  j                  |_        0 t        | j                  d       }|D ]  }t        |d      s|j                  } | j                  j                  d       | j                  t        | d|       t!        | d|       t#        | d|       | j                  j                  d       y )Nc                 $    t        | t               S )N)r*   r   )r;   s    r0   <lambda>z"_call_setup_hook.<locals>.<lambda>b   s    ZP[=\9\r<   )key
experiment	pre_setupsetupr$   
post_setup)r"   fnlightning_modulemodulesr*   r	   r   root_device_devicesortedr9   hasattrrA   barrierr7   r8   r6   _call_lightning_module_hook)r   rF   moduler9   r;   _s         r0   _call_setup_hookrQ   V   s    =='''			B **224f56$--99FN 5 W__*\]G 6<(!!A  [)%'C'4;\*r<   c                    t        d| j                        r/| j                  j                         5  t	        | d       d d d        t        d| j                        rv| j                  j                         5  | j                  j                         5  | j                  j                         5  t	        | d       d d d        d d d        d d d        y y # 1 sw Y   xY w# 1 sw Y   'xY w# 1 sw Y   +xY w# 1 sw Y   y xY w)Nconfigure_sharded_modelconfigure_model)r   rG   r   model_sharded_contextrN   tensor_init_contextprecision_pluginmodule_init_context)r   s    r0   _call_configure_modelrY   s   s    .0H0HI335'1JK 6
 &(@(@A002224$$88:'1BC ; 5 32 B 65 ;: 54 32sG   C6C7C+,C9C+C7CC($C++C4	0C77D c                    | j                   j                  J | j                   j                  }| j                  t        | d|       t	        | d|       t        | d|       d | j                  _        d | j                  _        | j                  D ]  }|j                  d        | j                  j                          y )NteardownrD   success)r"   rF   r7   r8   r6   rN   rG   _current_fx_name_metric_attributesr9   r:   profilerdescribe)r   rF   r;   s      r0   r   r      s    =='''			B%'2F*B72>04G-26G/ //	" " r<   )	pl_module	hook_namera   zpl.LightningModulec                   t         j                  | j                  j                   d|        |xs | j                  }|t        d      t        ||      }t        |      sy |j                  }||_        | j                  j                  d|j                  j                   d|       5   ||i |}d d d        ||_        S # 1 sw Y   xY w)Nz!: calling lightning module hook: z3No `LightningModule` is available to call hooks on.z[LightningModule].)logdebug	__class____name__rG   	TypeErrorgetattrcallabler]   r_   profile)r   rb   ra   r   r   rF   prev_fx_nameoutputs           r0   rN   rN      s     II""++,,Mi[YZ5W55IMNN	I	&BB<--L!*I				!	!$5i6I6I6R6R5SSTU^T_"`	aT$V$ 
b ".IM 
b	as   ,	CCc                    t         j                  | j                  j                   d|        | j                  t        d      t        | j                  |      }t        |      rQ| j                  j                  d| j                  j                  j                   d|       5   ||i |cd d d        S y # 1 sw Y   y xY w)Nz%: calling lightning datamodule hook: z7No `LightningDataModule` is available to call hooks on.z[LightningDataModule]rd   )
re   rf   rg   rh   r7   ri   rj   rk   r_   rl   )r   rb   r   r   rF   s        r0   r8   r8      s     II""++,,QR[Q\]^!QRR	##Y	/B|%%(=g>P>P>Z>Z>c>c=ddefoep&qrt&v& sr ss   &B99C)monitoring_callbacksrp   c                x   t         j                  | j                  j                   d|        | j                  }|r|j
                  }||_        | j                  }|du r'|D cg c]  }t        |t        t        f      s| }}n*|du r&|D cg c]  }t        |t        t        f      r| }}|D ]e  }	t        |	|      }
t        |
      s| j                  j                  d|	j                   d|       5   |
| | j                  g|i | d d d        g |r|_        y y c c}w c c}w # 1 sw Y   xY w)Nz: calling callback hook: TF
[Callback]rd   )re   rf   rg   rh   rG   r]   	callbacksr*   r   r
   rj   rk   r_   rl   	state_key)r   rb   rp   r   r   ra   rm   rs   cbcallbackrF   s              r0   r6   r6      s7    II""++,,Ei[QR((I 11%.	"!!It#"+[)Bz"}j>Y/ZR)	[		&"+_)B:b=R\B]3^R)	_Xy)B<!!))Jx7I7I6J!I;*WX7G44FtFvF YX  %1	"  \_
 YXs$   $D& D&D++D+9D00D9	c                 p    i }| j                   D ]$  }|j                         }|s|||j                  <   & |S )zzCalled when saving a model checkpoint, calls and returns every callback's `state_dict`, keyed by
    `Callback.state_key`.)rs   
state_dictrt   )r   callback_state_dictsrv   rx   s       r0   _call_callbacks_state_dictrz      sD     %%((*
7A !3!34 &  r<   
checkpointc                 2   | j                   }|r|j                  }d|_        | j                  D ]Q  }| j                  j	                  d|j
                   d      5  |j                  | | j                   |       ddd       S |r|_        yy# 1 sw Y   hxY w)zXCalled when saving a model checkpoint, calls every callback's `on_save_checkpoint` hook.on_save_checkpointrr   z.on_save_checkpointN)rG   r]   rs   r_   rl   rt   r}   )r   r{   ra   rm   rv   s        r0   "_call_callbacks_on_save_checkpointr~      s    ((I 11%9	"%%%%
83E3E2FFY&Z[''1I1I:V \[ & %1	"  \[s   BB	c                 T   | j                   }|r|j                  }d|_        |j                  d      }|yt        |d         t        d      k  }| j                  D ch c]  }|r|j
                  n|j                   }}|j                         |z
  }|rt        dt        |       d       | j                  D ]Q  }	| j                  j                  d|	j                   d	      5  |	j                  | | j                   |       ddd       S |r|_        yyc c}w # 1 sw Y   mxY w)
zCalled when loading a model checkpoint.

    Calls every callback's `on_load_checkpoint` hook. We have a dedicated function for this rather than using
    `_call_callback_hooks` because we have special logic for getting callback_states.

    on_load_checkpointrs   Nzpytorch-lightning_versionz1.5.0devzBe aware that when using `ckpt_path`, callbacks used to create the checkpoint need to be provided during `Trainer` instantiation. Please add the following callbacks: rd   rr   z.on_load_checkpoint)rG   r]   getr   rs   _legacy_state_keyrt   keysr   listr_   rl   r   )
r   r{   ra   rm   callback_statesis_legacy_ckptru   current_callbacks_keys
differencerv   s
             r0   "_call_callbacks_on_load_checkpointr     s7    ((I 11%9	">Hnn[>YOZ(CDEPZH[[Nahararsar[]nb22",,Vars %%'*@@J4484D3EQH	
 %%%%
83E3E2FFY&Z[''1I1I:V \[ & %1	"  t \[s   !D&DD'	c                     |j                  d      }|y| j                  D ]V  }|j                  |j                  |j                  |j                              }|s;t	        |      }|j                  |       X y)zQCalled when loading a model checkpoint, calls every callback's `load_state_dict`.rs   N)r   rs   rt   r   r   load_state_dict)r   r{   r   rv   r"   s        r0   _call_callbacks_load_state_dictr   *  sl    >Hnn[>YO%%##H$6$68K8KHLfLf8ghUOE$$U+	 &r<   c                    t         j                  | j                  j                   d|        | j                  }|j
                  }||_        t        | j                  |      }t        |      sy | j                  j                  d| j                  j                  j                   d|       5   ||i |}d d d        ||_        S # 1 sw Y   xY w)Nz: calling strategy hook: z
[Strategy]rd   )re   rf   rg   rh   rG   r]   rj   r   rk   r_   rl   )r   rb   r   r   ra   rm   rF   rn   s           r0   _call_strategy_hookr   8  s     II""++,,Ei[QR((I--L!*I	!!9	-BB<				!	!Jw/?/?/I/I/R/R.SSTU^T_"`	aT$V$ 
b ".IM 
b	as   /	C		C)r   r   r   N)6loggingr&   r,   copyr   typingr   r   r   r   packaging.versionr   pytorch_lightningpl-lightning_fabric.utilities.device_dtype_mixinr	   pytorch_lightning.callbacksr
   r   pytorch_lightning.loggersr   &pytorch_lightning.strategies.launchersr   5pytorch_lightning.trainer.connectors.signal_connectorr    pytorch_lightning.trainer.statesr   &pytorch_lightning.utilities.exceptionsr   )pytorch_lightning.utilities.model_helpersr   %pytorch_lightning.utilities.rank_zeror   r   	getLoggerrh   re   r1   r.   r)   rQ   rY   r   strrN   r8   boolr6   dictrz   r~   r   r   r    r<   r0   <module>r      s     
  1 1 %  Q A 1 L U : F C Pg!% %( %SV %be %jm %P" " "4 "+:D" 6 15	  ,-	
  	<  	
 	, ,0	222 2 #4.	2
 2 
2@   c4i  2 2$sTWx. 2]a 2 !2 !2$sTWx. !2]a !2H,\ ,tCQTH~ ,Z^ ,  	
 	r<   