a
    d                     @   s   d dl Z d dlZd dlZd dlZd dlmZmZmZmZm	Z	 d dl
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 d dlmZ e eZG dd	 d	eZeeeef d
ddZeedddZ dS )    N)AnyDictListOptionalUnion)_check_cuda_matmul_precisionnum_cuda_devices_parse_gpu_ids)_DEVICE)Accelerator)MisconfigurationExceptionc                   @   s   e Zd ZdZejddddZddddd	Zee	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	  dddZeee	 eej dddZee	dddZeedddZeeddddZdS )CUDAAcceleratorz$Accelerator for NVIDIA CUDA devices.Ndevicereturnc                 C   s2   |j dkrtd| dt| tj| dS )zs
        Raises:
            MisconfigurationException:
                If the selected device is not GPU.
        cudazDevice should be GPU, got z insteadN)typer   r   torchr   Z
set_deviceselfr    r   l/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/accelerators/cuda.pysetup_device#   s    
zCUDAAccelerator.setup_devicez
pl.Trainer)trainerr   c                 C   s   |  |j tj  d S N)set_nvidia_flags
local_rankr   r   empty_cache)r   r   r   r   r   setup.   s    zCUDAAccelerator.setup)r   r   c                 C   sL   dt jd< ddd tt D }t d|}td|  d| d	 d S )
NZ
PCI_BUS_IDZCUDA_DEVICE_ORDER,c                 s   s   | ]}t |V  qd S r   str.0xr   r   r   	<genexpr>8       z3CUDAAccelerator.set_nvidia_flags.<locals>.<genexpr>CUDA_VISIBLE_DEVICESzLOCAL_RANK: z - CUDA_VISIBLE_DEVICES: [])osenvironjoinranger   getenv_loginfo)r   Zall_gpu_idsdevicesr   r   r   r   4   s    
z CUDAAccelerator.set_nvidia_flagsc                 C   s   t j|S )a'  Gets stats for the given GPU device.

        Args:
            device: GPU device for which to get stats

        Returns:
            A dictionary mapping the metrics to their values.

        Raises:
            FileNotFoundError:
                If nvidia-smi installation not found
        )r   r   Zmemory_statsr   r   r   r   get_device_stats<   s    z CUDAAccelerator.get_device_stats)r   c                 C   s   t j  d S r   )r   r   r   )r   r   r   r   teardownK   s    zCUDAAccelerator.teardown)r1   r   c                 C   s   t | ddS )z!Accelerator device parsing logic.T)Zinclude_cudar	   r1   r   r   r   parse_devicesO   s    zCUDAAccelerator.parse_devicesc                 C   s   dd | D S )z*Gets parallel devices for the Accelerator.c                 S   s   g | ]}t d |qS )r   )r   r   r$   ir   r   r   
<listcomp>W   r'   z8CUDAAccelerator.get_parallel_devices.<locals>.<listcomp>r   r4   r   r   r   get_parallel_devicesT   s    z$CUDAAccelerator.get_parallel_devicesc                   C   s   t  S )z!Get the devices when set to auto.r   r   r   r   r   auto_device_countY   s    z!CUDAAccelerator.auto_device_countc                   C   s
   t  dkS )Nr   r:   r   r   r   r   is_available^   s    zCUDAAccelerator.is_available)accelerator_registryr   c                 C   s   |j d| | jj d d S )Nr   )description)register	__class____name__)clsr=   r   r   r   register_acceleratorsb   s
    z%CUDAAccelerator.register_accelerators)rA   
__module____qualname____doc__r   r   r   r   staticmethodintr   r   r   r"   r   r2   r3   r   r   r   r5   r9   r;   boolr<   classmethodrC   r   r   r   r   r       s"   (r   r   c                    s   t d}|du rtdg d}dd |D }d|}tj| }t|}tj	|d| d	d
| gdddd}t
tddd |j } fdd|dD }	dd t||	D }
|
S )a.  Get GPU stats including memory, fan speed, and temperature from nvidia-smi.

    Args:
        device: GPU device for which to get stats

    Returns:
        A dictionary mapping the metrics to their values.

    Raises:
        FileNotFoundError:
            If nvidia-smi installation not found
    z
nvidia-smiNznvidia-smi: command not found))zutilization.gpu%)zmemory.usedMB)zmemory.freerL   )zutilization.memoryrK   )z	fan.speedrK   )ztemperature.gpu   °C)ztemperature.memoryrM   c                 S   s   g | ]\}}|qS r   r   )r$   k_r   r   r   r8      r'   z(get_nvidia_gpu_stats.<locals>.<listcomp>r    z--query-gpu=z--format=csv,nounits,noheaderz--id=zutf-8T)encodingcapture_outputcheck)r%   r   c                 S   s$   z
t | W S  ty   Y dS 0 d S )Ng        )float
ValueError)r%   r   r   r   	_to_float   s    
z'get_nvidia_gpu_stats.<locals>._to_floatc                    s   g | ]} |qS r   r   r#   rU   r   r   r8      r'   z, c                 S   s&   i | ]\\}}}| d | d|qS )z ()r   )r$   r%   unitstatr   r   r   
<dictcomp>   r'   z(get_nvidia_gpu_stats.<locals>.<dictcomp>)shutilwhichFileNotFoundErrorr,   r   _utilsZ_get_device_index_get_gpu_id
subprocessrunr"   rS   stdoutstripsplitzip)r   Znvidia_smi_pathZgpu_stat_metricsZgpu_stat_keysZ	gpu_queryindexZgpu_idresultsstatsZ	gpu_statsr   rV   r   get_nvidia_gpu_statsk   s&    
	

rj   )	device_idr   c                 C   s:   d dd tt D }tjd|dd}||   S )zGet the unmasked real GPU IDs.r    c                 s   s   | ]}t |V  qd S r   r!   r6   r   r   r   r&      r'   z_get_gpu_id.<locals>.<genexpr>r(   )default)r,   r-   r   r*   r.   rd   rc   )rk   rl   Zcuda_visible_devicesr   r   r   r_      s    r_   )!loggingr*   r[   r`   typingr   r   r   r   r   r   Zpytorch_lightningplZ"lightning_fabric.accelerators.cudar   r   Z(lightning_fabric.utilities.device_parserr
   Z lightning_fabric.utilities.typesr   Z*pytorch_lightning.accelerators.acceleratorr   Z&pytorch_lightning.utilities.exceptionsr   	getLoggerrA   r/   r   r"   rS   rj   rH   r_   r   r   r   r   <module>   s   
K2