Ë
      çi+  ã                   óT   — d Z ddlmZ ddlmZ ddlmZ ddlZddl	m
Z
  G d„ d«      Zy)	z'Base class used to build new callbacks.é    )ÚAny)ÚTensor)Ú	OptimizerN)ÚSTEP_OUTPUTc                   ó  — e Zd ZdZedefd„«       Zeded    fd„«       Zde	defd„Z
ddd	d
deddfd„Zddd	d
deddfd„Zd=d„Zd=d„Zd=d„Zd=d„Zddd	d
de	deddf
d„Zddd	d
dede	deddfd„Zd=d„Zd=d„Zd=d„Zd=d„Zd=d„Zd=d„Zd=d„Zd=d„Z	 d>ddd	d
de	ded eddfd!„Z	 d>ddd	d
dede	ded eddfd"„Z	 d>ddd	d
de	ded eddfd#„Z	 d>ddd	d
dede	ded eddfd$„Z 	 d>ddd	d
de	ded eddfd%„Z!	 d>ddd	d
de	de	ded eddfd&„Z"d=d'„Z#d=d(„Z$d=d)„Z%d=d*„Z&d=d+„Z'd=d,„Z(d=d-„Z)d=d.„Z*ddd	d
d/e+ddfd0„Z,de-ee	f   fd1„Z.d2e-ee	f   ddfd3„Z/ddd	d
d4e-ee	f   ddfd5„Z0ddd	d
d4e-ee	f   ddfd6„Z1ddd	d
d7e2ddfd8„Z3d=d9„Z4ddd	d
d:e5ddfd;„Z6ddd	d
d:e5ddfd<„Z7y)?ÚCallbackzvAbstract base class used to build new callbacks.

    Subclass this class and override any of the relevant hooks

    Úreturnc                 ó.   — | j                   j                  S )au  Identifier for the state of the callback.

        Used to store and retrieve a callback's state from the checkpoint dictionary by
        ``checkpoint["callbacks"][state_key]``. Implementations of a callback need to provide a unique state key if 1)
        the callback has state and 2) it is desired to maintain the state of multiple instances of that callback.

        )Ú	__class__Ú__qualname__©Úselfs    úy/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pytorch_lightning/callbacks/callback.pyÚ	state_keyzCallback.state_key    s   € ð ~‰~×*Ñ*Ð*ó    c                 ó   — t        | «      S )z7State key for checkpoints saved prior to version 1.5.0.)Útyper   s    r   Ú_legacy_state_keyzCallback._legacy_state_key+   s   € ô D‹zÐr   Úkwargsc                 óH   — | j                   j                  › t        |«      › S )zÿFormats a set of key-value pairs into a state key string with the callback class name prefixed. Useful for
        defining a :attr:`state_key`.

        Args:
            **kwargs: A set of key-value pairs. Must be serializable to :class:`str`.

        )r   r   Úrepr)r   r   s     r   Ú_generate_state_keyzCallback._generate_state_key0   s"   € ð —.‘.×-Ñ-Ð.¬t°F«|¨nÐ=Ð=r   Útrainerú
pl.TrainerÚ	pl_moduleúpl.LightningModuleÚstageNc                  ó   — y)z9Called when fit, validate, test, predict, or tune begins.N© ©r   r   r   r   s       r   ÚsetupzCallback.setup:   ó    r   c                  ó   — y)z7Called when fit, validate, test, predict, or tune ends.Nr   r    s       r   ÚteardownzCallback.teardown=   r"   r   c                  ó   — y)zCalled when fit begins.Nr   ©r   r   r   s      r   Úon_fit_startzCallback.on_fit_start@   r"   r   c                  ó   — y)zCalled when fit ends.Nr   r&   s      r   Ú
on_fit_endzCallback.on_fit_endC   r"   r   c                  ó   — y)z/Called when the validation sanity check starts.Nr   r&   s      r   Úon_sanity_check_startzCallback.on_sanity_check_startF   r"   r   c                  ó   — y)z-Called when the validation sanity check ends.Nr   r&   s      r   Úon_sanity_check_endzCallback.on_sanity_check_endI   r"   r   ÚbatchÚ	batch_idxc                  ó   — y)z#Called when the train batch begins.Nr   )r   r   r   r.   r/   s        r   Úon_train_batch_startzCallback.on_train_batch_startL   r"   r   Úoutputsc                  ó   — y)záCalled when the train batch ends.

        Note:
            The value ``outputs["loss"]`` here will be the normalized value w.r.t ``accumulate_grad_batches`` of the
            loss returned from ``training_step``.

        Nr   )r   r   r   r2   r.   r/   s         r   Úon_train_batch_endzCallback.on_train_batch_endQ   r"   r   c                  ó   — y)z#Called when the train epoch begins.Nr   r&   s      r   Úon_train_epoch_startzCallback.on_train_epoch_start\   r"   r   c                  ó   — y)a+  Called when the train epoch ends.

        To access all batch outputs at the end of the epoch, you can cache step outputs as an attribute of the
        :class:`pytorch_lightning.core.LightningModule` and access them in this hook:

        .. code-block:: python

            class MyLightningModule(L.LightningModule):
                def __init__(self):
                    super().__init__()
                    self.training_step_outputs = []

                def training_step(self):
                    loss = ...
                    self.training_step_outputs.append(loss)
                    return loss


            class MyCallback(L.Callback):
                def on_train_epoch_end(self, trainer, pl_module):
                    # do something with all training_step outputs, for example:
                    epoch_mean = torch.stack(pl_module.training_step_outputs).mean()
                    pl_module.log("training_epoch_mean", epoch_mean)
                    # free up the memory
                    pl_module.training_step_outputs.clear()

        Nr   r&   s      r   Úon_train_epoch_endzCallback.on_train_epoch_end_   r"   r   c                  ó   — y)z!Called when the val epoch begins.Nr   r&   s      r   Úon_validation_epoch_startz"Callback.on_validation_epoch_start|   r"   r   c                  ó   — y)zCalled when the val epoch ends.Nr   r&   s      r   Úon_validation_epoch_endz Callback.on_validation_epoch_end   r"   r   c                  ó   — y)z"Called when the test epoch begins.Nr   r&   s      r   Úon_test_epoch_startzCallback.on_test_epoch_start‚   r"   r   c                  ó   — y)z Called when the test epoch ends.Nr   r&   s      r   Úon_test_epoch_endzCallback.on_test_epoch_end…   r"   r   c                  ó   — y)z%Called when the predict epoch begins.Nr   r&   s      r   Úon_predict_epoch_startzCallback.on_predict_epoch_startˆ   r"   r   c                  ó   — y)z#Called when the predict epoch ends.Nr   r&   s      r   Úon_predict_epoch_endzCallback.on_predict_epoch_end‹   r"   r   Údataloader_idxc                  ó   — y)z(Called when the validation batch begins.Nr   ©r   r   r   r.   r/   rE   s         r   Úon_validation_batch_startz"Callback.on_validation_batch_startŽ   r"   r   c                  ó   — y)z&Called when the validation batch ends.Nr   ©r   r   r   r2   r.   r/   rE   s          r   Úon_validation_batch_endz Callback.on_validation_batch_end˜   r"   r   c                  ó   — y)z"Called when the test batch begins.Nr   rG   s         r   Úon_test_batch_startzCallback.on_test_batch_start£   r"   r   c                  ó   — y)z Called when the test batch ends.Nr   rJ   s          r   Úon_test_batch_endzCallback.on_test_batch_end­   r"   r   c                  ó   — y)z%Called when the predict batch begins.Nr   rG   s         r   Úon_predict_batch_startzCallback.on_predict_batch_start¸   r"   r   c                  ó   — y)z#Called when the predict batch ends.Nr   rJ   s          r   Úon_predict_batch_endzCallback.on_predict_batch_endÂ   r"   r   c                  ó   — y)zCalled when the train begins.Nr   r&   s      r   Úon_train_startzCallback.on_train_startÍ   r"   r   c                  ó   — y)zCalled when the train ends.Nr   r&   s      r   Úon_train_endzCallback.on_train_endÐ   r"   r   c                  ó   — y)z'Called when the validation loop begins.Nr   r&   s      r   Úon_validation_startzCallback.on_validation_startÓ   r"   r   c                  ó   — y)z%Called when the validation loop ends.Nr   r&   s      r   Úon_validation_endzCallback.on_validation_endÖ   r"   r   c                  ó   — y)zCalled when the test begins.Nr   r&   s      r   Úon_test_startzCallback.on_test_startÙ   r"   r   c                  ó   — y)zCalled when the test ends.Nr   r&   s      r   Úon_test_endzCallback.on_test_endÜ   r"   r   c                  ó   — y)zCalled when the predict begins.Nr   r&   s      r   Úon_predict_startzCallback.on_predict_startß   r"   r   c                  ó   — y)zCalled when predict ends.Nr   r&   s      r   Úon_predict_endzCallback.on_predict_endâ   r"   r   Ú	exceptionc                  ó   — y)zACalled when any trainer execution is interrupted by an exception.Nr   )r   r   r   rd   s       r   Úon_exceptionzCallback.on_exceptionå   r"   r   c                 ó   — i S )z¡Called when saving a checkpoint, implement to generate callback's ``state_dict``.

        Returns:
            A dictionary containing callback state.

        r   r   s    r   Ú
state_dictzCallback.state_dictè   s	   € ð ˆ	r   rh   c                  ó   — y)zÅCalled when loading a checkpoint, implement to reload callback state given callback's ``state_dict``.

        Args:
            state_dict: the callback state returned by ``state_dict``.

        Nr   )r   rh   s     r   Úload_state_dictzCallback.load_state_dictñ   s   € ð 	r   Ú
checkpointc                  ó   — y)a  Called when saving a checkpoint to give you a chance to store anything else you might want to save.

        Args:
            trainer: the current :class:`~pytorch_lightning.trainer.trainer.Trainer` instance.
            pl_module: the current :class:`~pytorch_lightning.core.LightningModule` instance.
            checkpoint: the checkpoint dictionary that will be saved.

        Nr   ©r   r   r   rk   s       r   Úon_save_checkpointzCallback.on_save_checkpointú   r"   r   c                  ó   — y)ai  Called when loading a model checkpoint, use to reload state.

        Args:
            trainer: the current :class:`~pytorch_lightning.trainer.trainer.Trainer` instance.
            pl_module: the current :class:`~pytorch_lightning.core.LightningModule` instance.
            checkpoint: the full checkpoint dictionary that got loaded by the Trainer.

        Nr   rm   s       r   Úon_load_checkpointzCallback.on_load_checkpoint  r"   r   Úlossc                  ó   — y)z"Called before ``loss.backward()``.Nr   )r   r   r   rq   s       r   Úon_before_backwardzCallback.on_before_backward  r"   r   c                  ó   — y)zCCalled after ``loss.backward()`` and before optimizers are stepped.Nr   r&   s      r   Úon_after_backwardzCallback.on_after_backward  r"   r   Ú	optimizerc                  ó   — y)z#Called before ``optimizer.step()``.Nr   ©r   r   r   rv   s       r   Úon_before_optimizer_stepz!Callback.on_before_optimizer_step  r"   r   c                  ó   — y)z(Called before ``optimizer.zero_grad()``.Nr   rx   s       r   Úon_before_zero_gradzCallback.on_before_zero_grad  r"   r   )r   r   r   r   r	   N)r   )8Ú__name__Ú
__module__r   Ú__doc__ÚpropertyÚstrr   r   r   r   r   r!   r$   r'   r)   r+   r-   Úintr1   r   r4   r6   r8   r:   r<   r>   r@   rB   rD   rH   rK   rM   rO   rQ   rS   rU   rW   rY   r[   r]   r_   ra   rc   ÚBaseExceptionrf   Údictrh   rj   rn   rp   r   rs   ru   r   ry   r{   r   r   r   r   r      sP  „ ñð ð+˜3ò +ó ð+ð ð 4¨
Ñ#3ò ó ðð>¨Cð >°Có >ðH˜\ð HÐ6Jð HÐSVð HÐ[_ó HðF ð FÐ9Mð FÐVYð FÐ^bó Fó&ó$ó>ó<ð2Ø#ð2Ø0Dð2ØMPð2Ø]`ð2à	ó2ð
	Ø#ð	Ø0Dð	ØOZð	Øcfð	Øsvð	à	ó	ó2óó:0ó.ó1ó/ó4ó2ð  ñ7àð7ð (ð7ð ð	7ð
 ð7ð ð7ð 
ó7ð"  ñ	5àð	5ð (ð	5ð ð		5ð
 ð	5ð ð	5ð ð	5ð 
ó	5ð"  ñ1àð1ð (ð1ð ð	1ð
 ð1ð ð1ð 
ó1ð"  ñ	/àð	/ð (ð	/ð ð		/ð
 ð	/ð ð	/ð ð	/ð 
ó	/ð"  ñ4àð4ð (ð4ð ð	4ð
 ð4ð ð4ð 
ó4ð"  ñ	2àð	2ð (ð	2ð ð		2ð
 ð	2ð ð	2ð ð	2ð 
ó	2ó,ó*ó6ó4ó+ó)ó.ó(ðP Lð PÐ=Qð PÐ^kð PÐptó Pð˜D  c ™Nó ð¨$¨s°C¨x©.ð ¸Tó ð
Ø#ð
Ø0Dð
ØRVÐWZÐ\_ÐW_ÑR`ð
à	ó
ð
Ø#ð
Ø0Dð
ØRVÐWZÐ\_ÐW_ÑR`ð
à	ó
ð1¨,ð 1ÐCWð 1Ð_eð 1Ðjnó 1óRð2Ø#ð2Ø0Dð2ØQZð2à	ó2ð
7¨<ð 7ÐDXð 7Ðenð 7Ðswô 7r   r   )r~   Útypingr   Útorchr   Útorch.optimr   Úpytorch_lightningÚplÚ!pytorch_lightning.utilities.typesr   r   r   r   r   Ú<module>rŠ      s%   ðñ /å å Ý !ã Ý 9÷E7ò E7r   