a
    dL                     @   sl   d dl mZmZmZmZmZ d dlZd dlmZm	Z	 d dlm
Z d dlmZ d dlmZ G dd deZ
dS )	    )AnyDictListOptionalUnionN)_parse_tpu_devices_XLA_AVAILABLE)TPUAccelerator)_DEVICE)Acceleratorc                       s   e Zd ZdZeedd fddZejddddZe	e
eef dd	d
ZddddZeeeeee f eeeee f  dddZeeeee f ee dddZeedddZeedddZee
ddddZ  ZS )r	   zAccelerator for TPU devices.N)argskwargsreturnc                    s&   t sttt t j|i | d S N)r   ModuleNotFoundErrorstrsuper__init__)selfr   r   	__class__ k/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/accelerators/tpu.pyr      s    zTPUAccelerator.__init__)devicer   c                 C   s   d S r   r   )r   r   r   r   r   setup_device    s    zTPUAccelerator.setup_devicec                 C   s>   ddl m  m} ||}|d }|d | }||d}|S )zGets stats for the given TPU device.

        Args:
            device: TPU device for which to get stats

        Returns:
            A dictionary mapping the metrics (free memory and peak memory) to their values.
        r   NZkb_freeZkb_total)zavg. free memory (MB)zavg. peak memory (MB))Ztorch_xla.core.xla_modelcoreZ	xla_modelZget_memory_info)r   r   ZxmZmemory_infoZfree_memoryZpeak_memoryZdevice_statsr   r   r   get_device_stats#   s    	
zTPUAccelerator.get_device_stats)r   c                 C   s   d S r   r   )r   r   r   r   teardown7   s    zTPUAccelerator.teardown)devicesr   c                 C   s   t | S )z!Accelerator device parsing logic.)r   r   r   r   r   parse_devices:   s    zTPUAccelerator.parse_devicesc                 C   s   t | trtt| S | S )z*Gets parallel devices for the Accelerator.)
isinstanceintlistranger   r   r   r   get_parallel_devices?   s    
z#TPUAccelerator.get_parallel_devicesc                   C   s   dS )z!Get the devices when set to auto.   r   r   r   r   r   auto_device_countF   s    z TPUAccelerator.auto_device_countc                   C   s   t  S r   )FabricTPUAcceleratoris_availabler   r   r   r   r)   K   s    zTPUAccelerator.is_available)accelerator_registryr   c                 C   s   |j d| | jjd d S )NZtpu)description)registerr   __name__)clsr*   r   r   r   register_acceleratorsO   s
    z$TPUAccelerator.register_accelerators)r-   
__module____qualname____doc__r   r   torchr   r   r
   r   r   r   r   staticmethodr   r"   r   r   r    r%   r'   boolr)   classmethodr/   __classcell__r   r   r   r   r	      s   0"r	   )typingr   r   r   r   r   r3   Z!lightning_fabric.accelerators.tpur   r   r	   r(   Z lightning_fabric.utilities.typesr
   Z*pytorch_lightning.accelerators.acceleratorr   r   r   r   r   <module>   s   