a
    dM                     @   s   d Z ddlZddlZddlmZmZ ddlmZmZ ddl	m
Z
mZmZ ejejddZeded	d	d
d
de
jejdddZeded	d
d	d	e
jejdddZdS )a  This file exports ONNX ops for opset 16.

Note [ONNX Operators that are added/updated in opset 16]

~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
https://github.com/onnx/onnx/blob/main/docs/Changelog.md#version-16-of-the-default-onnx-operator-set
New operators:
    GridSample https://github.com/onnx/onnx/pull/3557

Updated operators:
    Identity
    If
    LeakyRelu
    Loop
    PRelu
    RoiAlign
    Scan
    ScatterElements
    ScatterND
    Where
    GreaterOrEqual
    LessOrEqual
    N)GRID_SAMPLE_INTERPOLATION_MODESGRID_SAMPLE_PADDING_MODES)_type_utilssymbolic_helper)	_beartype	jit_utilsregistration   )Zopsetzaten::grid_samplervib)gc                 C   s^   t |dkrt dS dd t D | }dd t D | }| jd||t|||dS )N   z#GridSample with 5D volumetric inputc                 S   s   i | ]\}}||qS  r   .0kr
   r   r   d/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/onnx/symbolic_opset16.py
<dictcomp>9       z grid_sampler.<locals>.<dictcomp>c                 S   s   i | ]\}}||qS r   r   r   r   r   r   r   :   r   Z
GridSample)Zalign_corners_imode_spadding_mode_s)r   Z_get_tensor_rankZ_onnx_unsupportedr   itemsr   opint)r   inputZgridZ	mode_enumZpadding_mode_enumZalign_cornersr   r   r   r   r   grid_sampler+   s    
r   zaten::scatter_addc           
      C   s*  t  r| jd||||ddS tj|tjj}t |}t |}t|t|krnt 	dd| d| dS ||ks~d |v r| 
d|}| j
d	td
gt| d}	| 
d||	|}t |}t |r| j
d||||ddS tj||kr| j
d|tj| d}| j
d||||ddS d S )NZscattersrc)Zoverload_namescatter_addz	`index` (z0) should have the same dimensionality as `src` ()ZShapeConstantr   )Zvalue_tSliceZScatterElementsadd)Zaxis_iZreduction_sZCast)Zto_i)r   Zis_caffe2_aten_fallbackatr   ZJitScalarType
from_value	UNDEFINEDZ_get_tensor_sizeslenZ_unimplementedr   torchZtensorZ_maybe_get_scalarZ	_is_valueZ	onnx_type)
r   selfZdimindexr   Zsrc_typeZ	src_sizesZindex_sizesZadjusted_shapeZstartsr   r   r   r   E   sB    

	

r   )__doc__	functoolsr'   Ztorch.nn.functionalr   r   Z
torch.onnxr   r   Ztorch.onnx._internalr   r   r   partialZonnx_symbolicZ_onnx_symbolic
parse_argsZbeartypeZGraphContextr   r   r   r   r   r   <module>   s   