
      iL:                     n   d Z ddlZddlZddl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dlmZmZmZmZmZ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mZmZ ddlmZ ddl m!Z!m"Z" ddl#m$Z$ ddl%m&Z&m'Z' erddl(m)Z)  ejT                  e+      Z,dZ- edd      Z. edd      Z/ G d de!      Z0defdZ1y)z
MLflow Logger
-------------
    N)	Namespace)Mapping)Path)time)TYPE_CHECKINGAnyCallableLiteralOptionalUnion)RequirementCache)Tensor)override)_add_prefix_convert_params_flatten_dict)ModelCheckpoint)Loggerrank_zero_experiment)_scan_checkpoints)rank_zero_onlyrank_zero_warnMlflowClientzfile:zmlflow>=1.0.0mlflowzmlflow>=2.8.0c                   H    e Zd ZdZdZdd ej                  d      dddddddf
d	ed
ee   dee   dee	ee
f      dee   ded   dedee   dee   dee   f fdZeed&d              Zedee   fd       Zedee   fd       Zeedee	ee
f   ef   ddfd              Zeed'deeef   dee   ddfd              Zeed(deddfd              Zeedee   fd               Zeedee   fd!              Zeedee   fd"              Zed#e ddfd$       Z!d#e ddfd%Z" xZ#S ))MLFlowLoggera
  Log using `MLflow <https://mlflow.org>`_.

    Install it with pip:

    .. code-block:: bash

        pip install mlflow  # or mlflow-skinny

    .. code-block:: python

        from lightning.pytorch import Trainer
        from lightning.pytorch.loggers import MLFlowLogger

        mlf_logger = MLFlowLogger(experiment_name="lightning_logs", tracking_uri="file:./ml-runs")
        trainer = Trainer(logger=mlf_logger)

    Use the logger anywhere in your :class:`~lightning.pytorch.core.LightningModule` as follows:

    .. code-block:: python

        from lightning.pytorch import LightningModule


        class LitModel(LightningModule):
            def training_step(self, batch, batch_idx):
                # example
                self.logger.experiment.whatever_ml_flow_supports(...)

            def any_lightning_module_function_or_hook(self):
                self.logger.experiment.whatever_ml_flow_supports(...)

    Args:
        experiment_name: The name of the experiment.
        run_name: Name of the new run. The `run_name` is internally stored as a ``mlflow.runName`` tag.
            If the ``mlflow.runName`` tag has already been set in `tags`, the value is overridden by the `run_name`.
        tracking_uri: Address of local or remote tracking server.
            If not provided, defaults to `MLFLOW_TRACKING_URI` environment variable if set, otherwise it falls
            back to `file:<save_dir>`.
        tags: A dictionary tags for the experiment.
        save_dir: A path to a local directory where the MLflow runs get saved.
            Defaults to `./mlruns` if `tracking_uri` is not provided.
            Has no effect if `tracking_uri` is provided.
        log_model: Log checkpoints created by :class:`~lightning.pytorch.callbacks.model_checkpoint.ModelCheckpoint`
            as MLFlow artifacts.

            * if ``log_model == 'all'``, checkpoints are logged during training.
            * if ``log_model == True``, checkpoints are logged at the end of training, except when
              :paramref:`~lightning.pytorch.callbacks.Checkpoint.save_top_k` ``== -1``
              which also logs every checkpoint during training.
            * if ``log_model == False`` (default), no checkpoint is logged.

        prefix: A string to put at the beginning of metric keys.
        artifact_location: The location to store run artifacts. If not provided, the server picks an appropriate
            default.
        run_id: The run identifier of the experiment. If not provided, a new run is started.
        synchronous: Hints mlflow whether to block the execution for every logging call until complete where
            applicable. Requires mlflow >= 2.8.0

    Raises:
        ModuleNotFoundError:
            If required MLFlow package is not installed on the device.

    -lightning_logsNMLFLOW_TRACKING_URIz./mlrunsF experiment_namerun_nametracking_uritagssave_dir	log_model)TFallprefixartifact_locationrun_idsynchronousc                    t         st        t        t                     |
t        st        d      t        |           |s
t         | }|| _        d | _        || _	        || _
        |	| _        || _        || _        i | _        d | _        || _        || _        |
i nd|
i| _        d| _        ddlm}  ||      | _        y )Nz$`synchronous` requires mlflow>=2.8.0r,   Fr   r   )_MLFLOW_AVAILABLEModuleNotFoundErrorstr_MLFLOW_SYNCHRONOUS_AVAILABLEsuper__init__LOCAL_FILE_URI_PREFIX_experiment_name_experiment_id_tracking_uri	_run_name_run_idr%   
_log_model_logged_model_time_checkpoint_callback_prefix_artifact_location_log_batch_kwargs_initializedmlflow.trackingr   _mlflow_client)selfr"   r#   r$   r%   r&   r'   r)   r*   r+   r,   r   	__class__s               u/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/pytorch/loggers/mlflow.pyr3   zMLFlowLogger.__init__t   s     !%c*;&<=="+H%&LMM34XJ?L /-1)!	#46?C!"3'2':P[@\!0*<8    returnc                 f   ddl }| j                  r| j                  S |j                  | j                         | j
                  S| j                  j                  | j
                        }|j                  j                  | _	        d| _        | j                  S | j                  | j                  j                  | j                        }|!|j                  dk7  r|j                  | _	        nYt        j                  d| j                   d       | j                  j                  | j                  | j                         | _	        | j
                  | j"                  h| j$                  xs i | _        ddlm} || j$                  v r&t        j                  d	| d
| j"                   d       | j"                  | j$                  |<   t+               }| j                  j-                  | j                   || j$                              }|j                  j.                  | _        d| _        | j                  S )zActual MLflow object. To use MLflow features in your :class:`~lightning.pytorch.core.LightningModule` do the
        following.

        Example::

            self.logger.experiment.some_mlflow_function()

        r   NTdeletedzExperiment with name z not found. Creating it.)namer*   )MLFLOW_RUN_NAMEzThe tag z3 is found in tags. The value will be overridden by .)experiment_idr%   )r   r@   rB   set_tracking_urir7   r9   get_runinforM   r6   get_experiment_by_namer5   lifecycle_stagelogwarningcreate_experimentr>   r8   r%   mlflow.utils.mlflow_tagsrK   _get_resolve_tags
create_runr+   )rC   r   runexptrK   resolve_tagss         rE   
experimentzMLFlowLogger.experiment   s    	&&& 2 23<<#%%--dll;C"%(("8"8D $D&&&&&&==d>S>STDD$8$8I$E&*&8&8#3D4I4I3JJbcd&*&9&9&K&K..$BYBY 'L '# <<~~) IIO	D"dii/KK"?"33fgkgugufvvwx .2^^		/*,.L%%00t?R?RYefjfofoYp0qC88??DL """rF   c                 2    | j                   }| j                  S )zqCreate the experiment if it does not exist to get the run id.

        Returns:
            The run id.

        )r\   r9   rC   _s     rE   r+   zMLFlowLogger.run_id   s     OO||rF   c                 2    | j                   }| j                  S )zCreate the experiment if it does not exist to get the experiment id.

        Returns:
            The experiment id.

        )r\   r6   r^   s     rE   rM   zMLFlowLogger.experiment_id   s     OO"""rF   paramsc           
      \   t        |      }t        |      }ddlm} |j	                         D cg c]  \  }} ||t        |      d d        }}}t        dt        |      d      D ];  } | j                  j                  d| j                  |||dz    d| j                   = y c c}}w )Nr   )Param   )keyvalued   )r+   ra    )r   r   mlflow.entitiesrc   itemsr0   rangelenr\   	log_batchr+   r?   )rC   ra   rc   kvparams_listidxs          rE   log_hyperparamszMLFlowLogger.log_hyperparams   s     !(v&) EKLLNSNDAqu#a&#,7NS C,c2C%DOO%%xT[[SSVY\S\A]xaeawawx 3 Ts    B(metricsstepc           
      P   t         j                  dk(  sJ d       ddlm} t	        || j
                  | j                        }g }t        t               dz        }|j                         D ]  \  }}t        |t              rt        j                  d| d| d       3t        j                  dd	|      }||k7  rt!        d
| d| dt"               |}|j%                   |||||xs d               | j&                  j(                  d| j*                  |d| j,                   y )Nr   z-experiment tried to log from global_rank != 0)Metrici  z$Discarding metric with string value =rL   z[^a-zA-Z0-9_/. -]+r!   zVMLFlow only allows '_', '/', '.' and ' ' special characters in metric name. Replacing z with )category)re   rf   	timestamprt   )r+   rs   rh   )r   rankri   rv   r   r=   LOGGER_JOIN_CHARintr   rj   
isinstancer0   rS   rT   resubr   RuntimeWarningappendr\   rm   r+   r?   )	rC   rs   rt   rv   metrics_listtimestamp_msrn   ro   new_ks	            rE   log_metricszMLFlowLogger.log_metrics   s    ""a'X)XX'*gt||T5J5JK%'46D=)MMODAq!S!B1#QqcKLFF/Q7EEz""#F5'4+
 1ATXT]\] ^_ $ 	"!!eledNdNderF   statusc                 2   | j                   sy |dk(  rd}n|dk(  rd}n|dk(  rd}| j                  r| j                  | j                         | j                  j	                  | j
                        r'| j                  j                  | j
                  |       y y )NsuccessFINISHEDfailedFAILEDfinished)r@   r<   _scan_and_log_checkpointsr\   rO   r+   set_terminated)rC   r   s     rE   finalizezMLFlowLogger.finalize  s       YFxFz!F $$**4+D+DE??""4;;/OO**4;;? 0rF   c                 z    | j                   j                  t              r| j                   t        t              d S y)zThe root file directory in which MLflow experiments are saved.

        Return:
            Local path to the root experiment directory if the tracking uri is local.
            Otherwise returns `None`.

        N)r7   
startswithr4   rl   rC   s    rE   r&   zMLFlowLogger.save_dir$  s6     (()>?%%c*?&@&BCCrF   c                     | j                   S )zQGet the experiment id.

        Returns:
            The experiment id.

        )rM   r   s    rE   rJ   zMLFlowLogger.name2  s     !!!rF   c                     | j                   S )zCGet the run id.

        Returns:
            The run id.

        )r+   r   s    rE   versionzMLFlowLogger.version=  s     {{rF   checkpoint_callbackc                     | j                   dk(  s| j                   du r!|j                  dk(  r| j                  |       y | j                   du r|| _        y y )Nr(   T)r:   
save_top_kr   r<   )rC   r   s     rE   after_save_checkpointz"MLFlowLogger.after_save_checkpointH  sR     ??e#t$'>CVCaCaegCg**+>?__$(;D% %rF   c                 h   t        || j                        }|D ]m  \  }}}}t        |t              r|j	                         n|t        |      j                  dD ci c]  }t        ||      r|t        ||       c}d}||j                  k(  rddgndg}	t        |      j                  }
| j                  j                  | j                  ||
       t        j                         5 }t!        | dd      5 }t#        j$                  ||d       d d d        t!        | d	d      5 }|j'                  t)        |	             d d d        | j                  j+                  | j                  ||
       d d d        || j                  |<   p y c c}w # 1 sw Y   xY w# 1 sw Y   \xY w# 1 sw Y   9xY w)
N)monitormode	save_lastr   save_weights_only_every_n_train_steps_every_n_val_epochs)scoreoriginal_filename
Checkpointlatestbestz/metadata.yamlwF)default_flow_stylez/aliases.txt)r   r;   r}   r   itemr   rJ   hasattrgetattrbest_model_pathstemr\   log_artifactr9   tempfileTemporaryDirectoryopenyamldumpwriter0   log_artifacts)rC   r   checkpointstpstagrn   metadataaliasesartifact_pathtmp_dirtmp_file_metadatatmp_file_aliasess                 rE   r   z&MLFlowLogger._scan_and_log_checkpointsP  s   '(;T=T=TU (LAq!S &06%:%)!W\\ 2A6 w2A66	H& -.1D1T1T,Tx([cZdG !GLLM OO((q-H ,,.'WI^4c:>OIIh(9eT ; WI\2C8<L$**3w<8 9 --dllG]S / *+D##A&U (
4 ;: 98 /.sB    F
&F(6FF(&F/F(FF(F%!F((F1	)rG   r   N)r   )$__name__
__module____qualname____doc__r{   osgetenvr0   r   dictr   r
   boolr3   propertyr   r\   r+   rM   r   r   r   r   rr   r   floatr|   r   r   r&   rJ   r   r   r   r   __classcell__)rD   s   @rE   r   r   1   s   >@   0"&&/bii0E&F)-",16+/ $&*%9%9 3-%9 sm	%9
 tCH~&%9 3-%9 -.%9 %9 $C=%9 %9 d^%9N 0#  0#d    #x} # # yeDcNI,E&F y4 y  y f73:#6 fhsm fW[ f  f4 @s @4 @  @" 
(3- 
  
 "hsm "  " #    < <T < </+_ /+QU /+rF   r   rG   c                  h    ddl m}  t        | d      rddlm} |S t        | d      rddlm} |S d }|S )Nr   )contextr[   )r[   registryc                     | S r   rh   )r%   s    rE   <lambda>z#_get_resolve_tags.<locals>.<lambda>  s    DrF   )rA   r   r   mlflow.tracking.contextr[    mlflow.tracking.context.registry)r   r[   s     rE   rW   rW     s@    ' w'8  
*	%A  )rF   )2r   loggingr   r~   r   argparser   collections.abcr   pathlibr   r   typingr   r   r	   r
   r   r   r    lightning_utilities.core.importsr   torchr   typing_extensionsr   !lightning.fabric.utilities.loggerr   r   r   ,lightning.pytorch.callbacks.model_checkpointr    lightning.pytorch.loggers.loggerr   r   #lightning.pytorch.loggers.utilitiesr   %lightning.pytorch.utilities.rank_zeror   r   rA   r   	getLoggerr   rS   r4   r.   r1   r   rW   rh   rF   rE   <module>r      s   
  	 	   #   I I  =  & Y Y H I A P,g! $_h?  0( K N+6 N+b
8 rF   