a
    d                     @   sJ   d dl Z d dlmZ d dlmZ d dlmZ e eZ	G dd deZ
dS )    N)Dict)_XLA_AVAILABLE)Profilerc                       sf   e Zd Zh dZh dZdedd fddZeddd	d
ZeddddZ	eedddZ
  ZS )XLAProfiler>   	test_steppredict_stepvalidation_step>   r   Zbackwardr   r   Ztraining_step4#  N)portreturnc                    s<   t sttt t jddd || _i | _i | _d| _dS )a*  XLA Profiler will help you debug and optimize training workload performance for your models using Cloud
        TPU performance tools.

        Args:
            port: the port to start the profiler server on. An exception is
                raised if the provided port is invalid or busy.
        N)dirpathfilenameF)	r   ModuleNotFoundErrorstrsuper__init__r
   _recording_map_step_recoding_map_start_trace)selfr
   	__class__ h/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/profilers/xla.pyr   "   s    zXLAProfiler.__init__)action_namer   c                 C   s   dd l m  m} |dd | jv r| js@|| j| _d| _|dd | j	v rn| 
|}|j||d}n
||}|  || j|< d S )Nr   .T)Zstep_num)Ztorch_xla.debug.profilerdebugZprofilersplitRECORD_FUNCTIONSr   Zstart_serverr
   serverSTEP_FUNCTIONS_get_step_numZ	StepTraceZTrace	__enter__r   )r   r   ZxpstepZ	recordingr   r   r   start2   s    

zXLAProfiler.startc                 C   s*   || j v r&| j | d d d  | j |= d S )N)r   __exit__r   r   r   r   r   stopD   s    
zXLAProfiler.stopc                 C   s2   || j vrd| j |< n| j |  d7  < | j | S )N   )r   r'   r   r   r   r"   I   s    
zXLAProfiler._get_step_num)r	   )__name__
__module____qualname__r!   r   intr   r   r%   r(   r"   __classcell__r   r   r   r   r      s   r   )loggingtypingr   Z!lightning_fabric.accelerators.tpur   Z$pytorch_lightning.profilers.profilerr   	getLoggerr*   logr   r   r   r   r   <module>   s
   
