Ë
    [^(h%  ã                  ó¾  — d Z ddlmZ ddlZddlZddlmZmZmZ ddl	m
Z
 ddlmZmZ g d¢Z ej                  ej                   d¬	«      Z ed
«       ej$                  d«      dd„«       «       Z ed«      d dd„«       Z ed«      d dd„«       Z ed«       ej,                  d«       ej$                  dd«      dd„«       «       «       Z ed«       ej$                  ddddddddd«	      	 	 dd„«       «       Z ed«      dd„«       Z ed«       ej$                  dddddddd«      	 	 	 	 	 d!	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d"d„«       «       Z	 	 	 	 	 	 d#d„Z	 	 	 	 	 	 	 	 d$d„Zy)%a&  This file exports ONNX ops for opset 14.

Note [ONNX operators that are added/updated in opset 14]
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
New operators:
    HardSwish, Trilu

Updated operators:
    Reshape
    Add, Sub, Mul, Div
    GRU, LSTM, RNN
    BatchNorm, Cumsum, Relu
é    )ÚannotationsN)Ú
_constantsÚ_type_utilsÚsymbolic_helper)ÚGLOBALS)Ú	jit_utilsÚregistration)Ú	hardswishÚtrilÚtriuÚreshapeÚ
batch_normÚquantized_hardswishÚscaled_dot_product_attentioné   )Úopsetzaten::hardswishÚvc                ó&   — | j                  d|«      S )NÚ	HardSwish©Úop)ÚgÚselfs     úY/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/onnx/symbolic_opset14.pyr
   r
   *   s   € ð �4‰4�˜TÓ"Ð"ó    z
aten::trilc                ó,   — | j                  d||d¬«      S )NÚTrilur   ©Úupper_ir   ©r   r   ÚdiagonalÚouts       r   r   r   0   ó   € à�4‰4�˜˜x°ˆ4Ó3Ð3r   z
aten::triuc                ó,   — | j                  d||d¬«      S )Nr   é   r   r   r    s       r   r   r   5   r#   r   zaten::reshapeTc                ó4   — t        j                  | ||d¬«      S )Nr   )Ú	allowzero)r   Ú_reshape_helper)r   r   Úshapes      r   r   r   :   s   € ô ×*Ñ*¨1¨d°EÀQÔGÐGr   zaten::batch_normÚiÚfc
                ó   — t        j                  «       rFt        j                  |||||g«      s,t        j
                  dk  rt        j                  dddd|«      S t        j                  |d«       t        j                  | |||||«      \  }}}}| j                  d||||||d|z
  |sdnd|sdnd¬	«
      }
|s|
S |
\  }}}|j                  |j                  «       «       |j                  |j                  «       «       |S )
Né   ÚBatchNormalizationr   zaAll input tensors must have the same `dtype`. Turn off Autocast or export using opset version 15.r   r%   r   é   )Ú	epsilon_fÚ
momentum_fÚtraining_mode_iÚoutputs)ÚtorchÚis_autocast_enabledr   Úargs_have_same_dtyper   Úexport_onnx_opset_versionÚ _onnx_opset_unsupported_detailedÚcheck_training_modeÚ_batchnorm_helperr   ÚsetTypeÚtype)r   ÚinputÚweightÚbiasÚrunning_meanÚrunning_varÚtrainingÚmomentumÚepsÚcudnn_enabledr"   ÚresÚnew_running_meanÚnew_running_vars                 r   r   r   C   s"  € ô 	×!Ñ!Ô#Ü×4Ñ4Ø�F˜D ,°Ð<ô
ô ×-Ñ-°Ò2ä×?Ñ?Ø ØØðCàó
ð 	
ô ×'Ñ'¨°,Ô?Ü.=×.OÑ.OØ	ˆ5�&˜$ ¨kó/Ñ+€FˆD�, ð �$‰$ØØØØØØØØ�x‘<Ù!)™¨qÙ!‘ qð ó €Cñ Øˆ
à14Ñ.ˆÐ˜Ø× Ñ  ×!2Ñ!2Ó!4Ô5Ø×Ñ × 0Ñ 0Ó 2Ô3Øˆ
r   zquantized::hardswishc                ó€   — t        j                  | |«      \  }}}}t        | |«      }t        j                  | |||«      S ©N)r   Údequantize_helperr
   Úquantize_helper)r   ÚxÚop_scaleÚop_zero_pointÚ_Úoutputs         r   r   r   z   s>   € ä ×2Ñ2°1°aÓ8�J€A€qˆ!ˆQä�q˜!‹_€Fä×*Ñ*¨1¨f°hÀÓNÐNr   z"aten::scaled_dot_product_attentionÚbc	                óä  — |r|rt        j                  |«      sJ d«       ‚|rJ d«       ‚t        j                  |«      rt        | |«      }|rt        | ||«      }t        j                  |«      }	t        t        |	«      «      }
|
d   |
d   c|
d<   |
d<   | j                  d||
¬«      }| j                  d|| j                  d|«      «      }| j                  d|| j                  d|«      «      }| j                  d	||«      }t        j                  |«      r|}�net        j                  j                  |«      t        j                  j                  k(  r€| j                  d
t        j                  dg«      ¬«      }| j                  d
t        j                  t        d«       g«      ¬«      }| j                  d|||«      }| j                  d||«      }n«t        j                  j                  |«      t        j                  j                  t        j                  j                   t        j                  j"                  fv r| j                  d||«      }n+t%        dt        j                  j                  |«      › �«      ‚| j                  d|d¬«      }|dk7  rG| j                  d|| j                  d
t        j                  |t        j                  ¬«      ¬«      «      }| j                  d	||«      S )Nz6is_causal and attn_mask cannot be set at the same timezPconversion of scaled_dot_product_attention not implemented if enable_gqa is TrueéþÿÿÿéÿÿÿÿÚ	Transpose)Úperm_iÚMulÚSqrtÚMatMulÚConstantç        ©Úvalue_tÚinfÚWhereÚAddz Unsupported type for attn_mask: ÚSoftmax©Úaxis_ir   ÚDropout©Údtype)r   Ú_is_noneÚ_attention_scaleÚ_causal_attention_maskÚ_get_tensor_rankÚlistÚranger   r   ÚJitScalarTypeÚ
from_valueÚBOOLr4   ÚtensorÚfloatÚFLOATÚHALFÚBFLOAT16Ú
ValueError)r   ÚqueryÚkeyÚvalueÚ	attn_maskÚ	dropout_pÚ	is_causalÚscaleÚ
enable_gqaÚkey_shape_builtinÚkey_transposed_axesÚkey_transposedÚquery_scaledÚkey_transposed_scaledÚmul_qkÚ
mul_qk_addÚ
const_zeroÚconst_neg_infÚattn_weights                      r   r   r   ‡   s¨  € ñ ™y¬_×-EÑ-EÀiÔ-Pð Ø@óÐQñ ð ØZóˆ>ô ×Ñ Ô&Ü   EÓ*ˆáÜ*¨1¨e°SÓ9ˆ	ô
 (×8Ñ8¸Ó=ÐÜœuÐ%6Ó7Ó8Ðà˜BÑØ˜BÑð 5Ð˜ÑÐ0°Ñ4ð —T‘T˜+ sÐ3F�TÓG€Nð —4‘4˜˜u a§d¡d¨6°5Ó&9Ó:€LØŸD™D ¨¸¿¹¸VÀUÓ8KÓLÐØ�T‰T�(˜LÐ*?Ó@€Fä×Ñ 	Ô*ØŠ
ä×!Ñ!×,Ñ,¨YÓ7Ü×$Ñ$×)Ñ)ò	*ð —T‘T˜*¬e¯l©l¸C¸5Ó.A�TÓBˆ
ØŸ™˜Z´·±ÄÀeÃ¸}¸oÓ1N˜ÓOˆØ—D‘D˜ )¨Z¸ÓGˆ	Ø—T‘T˜% ¨Ó3‰
Ü	×	"Ñ	"×	-Ñ	-¨iÓ	8Ü×!Ñ!×'Ñ'Ü×!Ñ!×&Ñ&Ü×!Ñ!×*Ñ*ð=ñ 
ð
 —T‘T˜% ¨Ó3‰
äØ.¬{×/HÑ/H×/SÑ/SÐT]Ó/^Ð._Ð`ó
ð 	
ð —$‘$�y *°R�$Ó8€Kà�A‚~Ø—d‘dØØØ�D‰D�¤U§\¡\°)Ä5Ç;Á;Ô%OˆDÓPó
ˆð �4‰4�˜+ uÓ-Ð-r   c                óò  — | j                  d|«      }| j                  d|| j                  dt        j                  dgt        j                  ¬«      ¬«      | j                  dt        j                  t        j
                  gt        j                  ¬«      ¬«      «      }| j                  d|t        j                  j                  |«      j                  «       ¬«      }| j                  dt        j                  d	gt        j                  ¬«      ¬«      }| j                  d
|| j                  d|«      «      }| j                  d|t        j                  j                  |«      j                  «       ¬«      }|S )zºCalculate the scale factor for the attention result.

    Args:
        query: Tensor of shape [..., L, E]

    Returns:
        Scalar scale factor := 1 / math.sqrt(query.size(-1))
    ÚShapeÚSlicer[   rU   rf   r]   ÚCast)Úto_iç      ð?ÚDivrY   )r   r4   rq   Úint64r   Ú	INT64_MAXr   rn   ro   Ú	onnx_typerr   )r   rw   Úquery_shapeÚquery_shape_lastÚembedding_sizeÚ	const_oner}   s          r   ri   ri   Ô   s/  € ð —$‘$�w Ó&€KØ—t‘tØØØ	�‰ˆZ¤§¡¨r¨d¼%¿+¹+Ô!FˆÓGØ	�‰Ø¤§¡¬j×.BÑ.BÐ-CÌ5Ï;É;Ô Wð 	ó 	
ó	Ðð —T‘TØØÜ×&Ñ&×1Ñ1°%Ó8×BÑBÓDð ó €Nð
 —‘�Z¬¯©°s°eÄ5Ç;Á;Ô)O�ÓP€IØ�D‰D�˜	 1§4¡4¨°Ó#?Ó@€Eà�D‰DØØÜ×&Ñ&×1Ñ1°%Ó8×BÑBÓDð ó €Eð
 €Lr   c                ó:  — | j                  d|«      }| j                  d|«      }| j                  dt        j                  dgt        j                  ¬«      ¬«      }| j                  dt        j                  dgt        j                  ¬«      ¬«      }| j                  d|||«      }| j                  d|||«      }| j                  d||d	¬
«      }	| j                  dt        j                  dg«      ¬«      }
| j                  d|
|	«      }| j                  d|d	¬«      }| j                  dt        j                  dg«      ¬«      }| j                  dt        j                  t	        d«       g«      ¬«      }| j                  d| j                  d||«      ||«      }|S )až  Create a causal mask for the given query and key tensors.

    Equivalent to::
        mask = torch.ones(L, S, dtype=torch.bool).tril(diagonal=0)
        attn_mask = torch.zeros(L, S, dtype=torch.float)
        attn_mask = attn_mask.masked_fill(not mask, -float("inf"))

    Args:
        query: Tensor of shape [..., L, E]
        key: Tensor of shape [..., S, E]

    Returns:
        Tensor of shape [L, S]
    rŠ   r[   rU   rf   r]   rT   r‹   ÚConcatr   rc   rŽ   ÚExpandr   r   r\   r_   r`   ÚEqual)r   r4   rq   r�   rr   )r   rw   rx   r“   Ú	key_shapeÚlast_idxÚsecond_last_idxÚtarget_lengthÚsource_lengthÚsizer–   rz   r†   r‡   s                 r   rj   rj   ø   sW  € ð$ —$‘$�w Ó&€KØ—‘�W˜cÓ"€Ià�t‰t�J¬¯©°b°TÄÇÁÔ(MˆtÓN€HØ—d‘d˜:¬u¯|©|¸R¸DÌÏÉÔ/T�dÓU€OØ—D‘D˜ +¨ÀÓI€MØ—D‘D˜ )¨_¸hÓG€Mà�4‰4�˜-¨¸qˆ4ÓA€DØ—‘�Z¬¯©°s°eÓ)<�Ó=€IØ—‘�X˜y¨$Ó/€Ià—‘�W˜i°�Ó3€Ià—‘�j¬%¯,©,¸°uÓ*=�Ó>€JØ—D‘D˜¬U¯\©\¼EÀ%»L¸=¸/Ó-J�DÓK€MØ—‘Ø�—‘�g˜y¨*Ó5°}Àjó€Ið Ðr   )r   újit_utils.GraphContextrJ   )Nr\   FNF)r   r¡   rw   útorch._C.Valuerx   r¢   ry   r¢   rz   útorch._C.Value | Noner{   rr   r|   Úboolr}   r£   r~   r¤   )r   r¡   rw   r¢   Úreturnr¢   )r   r¡   rw   r¢   rx   r¢   r¥   r¢   )Ú__doc__Ú
__future__r   Ú	functoolsr4   Ú
torch.onnxr   r   r   Útorch.onnx._globalsr   Útorch.onnx._internalr   r	   Ú__all__ÚpartialÚonnx_symbolicÚ_onnx_symbolicÚ
parse_argsr
   r   r   Úquantized_argsr   r   r   r   ri   rj   © r   r   ú<module>r³      sE  ðñõ  #ã ã ß ?Ñ ?Ý 'ß 8ò€ð #�×"Ñ" <×#=Ñ#=ÀRÔH€ñ Ð!Ó"Ø€×Ñ˜CÓ ò#ó !ó #ð#ñ �Óó4ó ð4ñ �Óó4ó ð4ñ �Ó Ø€×Ñ Ó%Ø€×Ñ˜C Ó%òHó &ó &ó !ðHñ Ð"Ó#Ø€×Ñ˜C  c¨3°°S¸#¸sÀCÓHð2Øò2ó Ió $ð2ñj Ð&Ó'òOó (ðOñ Ð4Ó5Ø€×Ñ˜C  c¨3°°S¸#¸sÓCð (,ØØØ#'ØðH.ØðH.àðH.ð 
ðH.ð ð	H.ð
 %ðH.ð ðH.ð ðH.ð !ðH.ð òH.ó Dó 6ðH.ðV!Øð!Ø&4ð!àó!ðH%Øð%Ø&4ð%Ø;Ið%àô%r   