a
    d	                     @   s   d dl mZmZmZ d dlmZmZ d dlmZ d dl	m
Z
 d dlmZ erXd dlmZ ed Zed Zeeef ZG d	d
 d
e
ZdS )    )castOptionalUnion)get_argsLiteral)_HPU_AVAILABLE)PrecisionPlugin)MisconfigurationException)hmp)       )Z3216bf16c                   @   s2   e Zd ZdZdeeee ee eddddZdS )	HPUPrecisionPluginaI  Plugin that enables bfloat/half support on HPUs.

    Args:
        precision: The precision to use.
        opt_level: Choose optimization level for hmp.
        bf16_file_path: Path to bf16 ops list in hmp O1 mode.
        fp32_file_path: Path to fp32 ops list in hmp O1 mode.
        verbose: Enable verbose mode for hmp.
    O2NF)	precision	opt_levelbf16_file_pathfp32_file_pathverbosereturnc                 C   sj   t stdtttt }||vr:td|d| dttt|| _| jdv rft	j
||||d d S )Nz*HPU precision plugin requires HPU devices.z&`Trainer(accelerator='hpu', precision=z1)` is not supported. `precision` must be one of: .)r   r   )r   r   r   Z	isVerbose)r   r	   r   _PRECISION_INPUT_STR_PRECISION_INPUT_INT
ValueErrorr   strr   r
   convert)selfr   r   r   r   r   Zsupported_precision r   p/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/plugins/precision/hpu.py__init__)   s    
zHPUPrecisionPlugin.__init__)r   NNF)	__name__
__module____qualname____doc___PRECISION_INPUTr   r   boolr    r   r   r   r   r      s       r   N)typingr   r   r   Ztyping_extensionsr   r   Z"pytorch_lightning.accelerators.hpur   Z4pytorch_lightning.plugins.precision.precision_pluginr   Z&pytorch_lightning.utilities.exceptionsr	   Zhabana_frameworks.torch.hpexr
   r   r   r%   r   r   r   r   r   <module>   s   