a
    du                     @   s8  d dl Z d dlZd dlZd dlmZ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 G dd deZd	Zee	eedd
ddZe	e	dddZeedddZedZedddZeeeeee f  eeeee f  dddZeedddZ eeeee f dddZ!dS )    N)ProcessQueue)AnyCallableDictListOptionalUnion)ModuleAvailableCache)Accelerator)_check_data_typec                       s   e Zd ZdZeedd fddZejd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jdded	ddZeeddddZ  ZS )TPUAcceleratorzAccelerator for TPU devices.Nargskwargsreturnc                    s&   t sttt t j|i | d S N)_XLA_AVAILABLEModuleNotFoundErrorstrsuper__init__)selfr   r   	__class__ j/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lightning_fabric/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_devicer   c                 C   s   d S r   r   )r   r   r   r   teardown&   s    zTPUAccelerator.teardowndevicesr   c                 C   s   t | S )z!Accelerator device parsing logic.)_parse_tpu_devices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_count5   s    z TPUAccelerator.auto_device_count   )maxsizec                   C   s   t tot S r   )boolr   _is_device_tpur   r   r   r   is_available:   s    zTPUAccelerator.is_available)accelerator_registryr   c                 C   s   |j d| | jjd d S )NZtpu)description)registerr   __name__)clsr2   r   r   r   register_accelerators@   s
    z$TPUAccelerator.register_accelerators)r5   
__module____qualname____doc__r   r   torchr   r   r    staticmethodr	   r'   r   r   r   r%   r*   r,   	functools	lru_cacher/   r1   classmethodr   r7   __classcell__r   r   r   r   r      s   0"
r   <   )queuefuncr   r   r   c                 O   sB   z|  ||i | W n$ ty<   t  |  d  Y n0 d S r   )put	Exception	traceback	print_exc)rB   rC   r   r   r   r   r   _inner_fM   s
    rH   )rC   r   c                    s,   t  tttttf d fdd}|S )Nr   c                     s^   t  }tt| g| R |d}|  |t z
| W S  tjyX   t	
  Y dS 0 d S )N)targetr   r   F)r   r   rH   startjoinTPU_CHECK_TIMEOUT
get_nowaitqEmptyrF   rG   )r   r   rB   procrC   r   r   wrapperV   s    

z_multi_process.<locals>.wrapper)r=   wrapsr   r	   r/   )rC   rR   r   rQ   r   _multi_processU   s     rT   r   c                  C   s4   t sdS ddlm  m}  |  dkp2t| dS )zCheck if TPU devices are available. Runs XLA device check within a separate process.

    Return:
        A boolean value indicating if TPU devices are available
    Fr   Nr-   ZTPU)r   torch_xla.core.xla_modelcore	xla_modelxrt_world_sizer/   Zget_xla_supported_devicesZxmr   r   r   r0   e   s    r0   Z	torch_xlac                  C   s*   t  sdS dd lm  m}  |  dkS )NFr   r-   )r   r1   rU   rV   rW   rX   rY   r   r   r   _tpu_distributed{   s    rZ   r!   c                 C   s2   t |  t| trt|  } t| s.td| S )a  
    Parses the TPU devices given in the format as accepted by the
    :class:`~pytorch_lightning.trainer.Trainer` and :class:`~lightning_fabric.Fabric`.

    Args:
        devices: An int of 1 or string '1' indicates that 1 core with multi-processing should be used
            An int 8 or string '8' indicates that all 8 cores with multi-processing should be used
            A list of ints or a strings containing a list of comma separated integers
            indicates the specific TPU core to use.

    Returns:
        A list of tpu_cores to be used or ``None`` if no TPU cores were requested

    Raises:
        TypeError:
            If TPU devices aren't 1, 8 or [<1-8>]
    z/`devices` can only be 1, 8 or [<1-8>] for TPUs.)r   r&   r   _parse_tpu_devices_strstrip_tpu_devices_valid	TypeErrorr$   r   r   r   r#      s    
r#   c                 C   sX   | dv rdS t | tttfrTt| dk}dt| d   koBdkn  }|oN|}|S dS )N)r-   r+   NTr-   r   r+   F)r&   r(   tuplesetlen)r"   Zhas_1_tpu_idxZis_valid_tpu_idxZis_valid_tpu_core_choicer   r   r   r]      s     r]   c                 C   s$   | dv rt | S dd | dD S )N)18c                 S   s$   g | ]}t |d krt| qS )r   )ra   r'   r\   ).0xr   r   r   
<listcomp>       z*_parse_tpu_devices_str.<locals>.<listcomp>,)r'   splitr$   r   r   r   r[      s    r[   )"r=   rB   rN   rF   multiprocessingr   r   typingr   r   r   r   r   r	   r;   Z lightning_utilities.core.importsr
   Z)lightning_fabric.accelerators.acceleratorr   Z(lightning_fabric.utilities.device_parserr   r   rL   rH   rT   r/   r0   r   rZ   r'   r   r#   r]   r[   r   r   r   r   <module>   s&    /2