o
    »i e: ã                   @   sŒ  d Z ddlZddlmZ ddlmZ ddlmZ ddlm	Z	 ddlm
Z
 ddlmZ dd	lmZ dd
lmZ ddlmZ ddlmZ ddlmZ ddlmZ dd„ Ze
 d¡dd„ ƒZe
 d¡dd„ ƒZe
 d¡dd„ ƒZdd„ ZdZdd„ Ze
 d¡d d!„ ƒZd"d#„ Ze
 d$¡d%d&„ ƒZe
 d'¡d(d)„ ƒZe
 d*¡d+d,„ ƒZe
 d-¡d.d/„ ƒZ e
 d0¡d1d2„ ƒZ!e
 d3¡d4d5„ ƒZ"e
 d6¡d7d8„ ƒZ#e
 d9¡d:d;„ ƒZ$e
 d<¡d=d>„ ƒZ%e
 d?¡d@dA„ ƒZ&e
 dB¡dCdD„ ƒZ'e
 dE¡dFdG„ ƒZ(dHdI„ Z)e
 dJ¡dKdL„ ƒZ*e
 dM¡dNdO„ ƒZ+e
 dP¡dQdR„ ƒZ,		d§dSdT„Z-dUdV„ Z.e
 dW¡dXdY„ ƒZ/e
 dZ¡d[d\„ ƒZ0e
 d]¡d^d_„ ƒZ1e
 d`¡dadb„ ƒZ2e
 dc¡ddde„ ƒZ3e
 df¡dgdh„ ƒZ4e
 di¡djdk„ ƒZ5e
 dl¡dmdn„ ƒZ6e
 do¡dpdq„ ƒZ7e
 dr¡dsdt„ ƒZ8e
 du¡dvdw„ ƒZ9e
 dx¡dydz„ ƒZ:e
 d{¡d|d}„ ƒZ;e
 d~¡dd€„ ƒZ<e
 d¡d‚dƒ„ ƒZ=e
 d„¡d…d†„ ƒZ>e
 d‡¡dˆd‰„ ƒZ?e
 dŠ¡d‹dŒ„ ƒZ@e
 d¡dŽd„ ƒZAe
 d¡d‘d’„ ƒZBe
 d“¡d”d•„ ƒZCe
 d–¡d—d˜„ ƒZDe
 d™¡dšd›„ ƒZEe
 dœ¡ddž„ ƒZFe
 dŸ¡d d¡„ ƒZGe
 d¢¡d£d¤„ ƒZHe
 d¥¡d¦d§„ ƒZIe
 d¨¡d©dª„ ƒZJe
 d«¡d¬d­„ ƒZKe
 d®¡d¯d°„ ƒZLe
 d±¡d²d³„ ƒZMe
 d´¡dµd¶„ ƒZNe
 d·¡d¸d¹„ ƒZOe
 dº¡d»d¼„ ƒZPe
 d½¡d¾d¿„ ƒZQe
 dÀ¡dÁdÂ„ ƒZRe
 dÃ¡dÄdÅ„ ƒZSe
 dÆ¡dÇdÈ„ ƒZTe
 dÉ¡dÊdË„ ƒZUe
 dÌ¡dÍdÎ„ ƒZVe
 dÏ¡dÐdÑ„ ƒZWe
 dÒ¡dÓdÔ„ ƒZXe
 dÕ¡dÖd×„ ƒZYe
 dØ¡dÙdÚ„ ƒZZe
 dÛ¡dÜdÝ„ ƒZ[e
 dÞ¡dßdà„ ƒZ\e
 dá¡dâdã„ ƒZ]e
 dä¡dådæ„ ƒZ^e
 dç¡dèdé„ ƒZ_e
 dê¡dëdì„ ƒZ`e
 dí¡dîdï„ ƒZae
 dð¡dñdò„ ƒZbe
 dó¡dôdõ„ ƒZce
 dö¡d÷dø„ ƒZde
 dù¡dúdû„ ƒZee
 dü¡dýdþ„ ƒZfe
 dÿ¡d d„ ƒZge
 d¡dd„ ƒZhe
 d¡dd„ ƒZie
 d¡d	d
„ ƒZje
 d¡dd„ ƒZke
 d¡dd„ ƒZle
 d¡dd„ ƒZme
 d¡dd„ ƒZne
 d¡dd„ ƒZoe
 d¡dd„ ƒZpe
 d¡dd„ ƒZqe
 d ¡d!d"„ ƒZrd#d$„ Zse
 d%¡e
 d&¡d'd(„ ƒƒZte
 d)¡d*d+„ ƒZue
 d,¡d-d.„ ƒZve
 d/¡d0d1„ ƒZwe
 d2¡d3d4„ ƒZxe
 d5¡d6d7„ ƒZye
 d8¡d9d:„ ƒZze
 d;¡d<d=„ ƒZ{e
 d>¡d?d@„ ƒZ|e
 dA¡dBdC„ ƒZ}e
 dD¡dEdF„ ƒZ~dGdH„ ZdIdJ„ Z€e
 dK¡dLdM„ ƒZe
 dN¡dOdP„ ƒZ‚e
 dQ¡dRdS„ ƒZƒe
 „dT¡ e
 „dU¡ e
 „dV¡ e
 „dW¡ e
 „dX¡ e
 „dY¡ e
 „dZ¡ e
 „d[¡ e
 „d\¡ e
 „d]¡ e
 d^¡d_d`„ ƒZ…e
 da¡dbdc„ ƒZ†ddde„ Z‡dfdg„ Zˆe
 dh¡didj„ ƒZ‰e
 dk¡dldm„ ƒZŠe
 dn¡dodp„ ƒZ‹e
 dq¡drds„ ƒZŒe
 dt¡dudv„ ƒZe
 dw¡dxdy„ ƒZŽe
 dz¡d{d|„ ƒZe
 d}¡e
 d~¡dd€„ ƒƒZe
 „d¡ e
 „d‚¡ e
 dƒ¡d„d…„ ƒZ‘e
 d†¡d‡dˆ„ ƒZ’e
 d‰¡dŠd‹„ ƒZ“e
 dŒ¡ddŽ„ ƒZ”e
 d¡dd‘„ ƒZ•e
 d’¡d“d”„ ƒZ–e
 d•¡d–d—„ ƒZ—e
 d˜¡d™dš„ ƒZ˜e
 d›¡dœd„ ƒZ™e
 dž¡dŸd „ ƒZše
 d¡¡d¢d£„ ƒZ›e
 d¤¡d¥d¦„ ƒZœdS (¨  z/Gradients for operators defined in math_ops.py.é    N)Úcompat)Úcontext)Úconstant_op)Údtypes)Úops)Útensor)Útensor_util)Ú	array_ops)Úgen_array_ops)Úgen_math_ops)Úmath_ops)Úspecial_math_opsc                 C   s   | t  |d¡ S )z;Divides `x / y` assuming `x, y >= 0`, treating `0 / 0 = 0`.é   )r   Úmaximum)ÚxÚy© r   úN/var/www/myenv/lib/python3.10/site-packages/tensorflow/python/ops/math_grad.pyÚ_safe_shape_div    s   r   ÚArgMaxc                 C   ó   ~ ~d d gS ©Nr   ©ÚopÚgradr   r   r   Ú_ArgMaxGrad%   ó   r   ÚArgMinc                 C   r   r   r   r   r   r   r   Ú_ArgMinGrad+   r   r   ÚEuclideanNormc                 C   sd   | j d }|  d¡s%t t | jd ¡| jd ¡}t ||¡}t ||¡}t | jd || ¡dfS )zGradient for EuclideanNorm.r   Ú	keep_dimsr   N)	ÚoutputsÚget_attrr   Úreduced_shaper	   ÚshapeÚinputsÚreshapeÚtruediv)r   r   ÚoutputÚoutput_shape_kept_dimsr   r   r   Ú_EuclideanNormGrad1   s   

ÿr*   c                 C   sš  t  ¡ st| tjƒrt|tjƒrt|tjƒs2t | ¡}t |¡}t ||¡\}}||df||dffS |  	¡ }| 	¡ }| 	¡ }	|du sNd|v sN|du sNd|v rntj
| dd}tj
|dd}t ||¡\}}||df||dffS ||	k}
||	k}t ¡ }z|j||f \}}|||
f|||ffW S  tyÌ   t ||¡\}}tt |¡ƒ}|dusªJ ‚tt |¡ƒ}|dus·J ‚||f|j||f< |||
f|||ff Y S w )aÄ  Optimized version of `broadcast_gradient_args` that caches results.

  This implementation avoids creating `broadcast_gradient_args` ops in the case
  that the input shapes are fully defined, and provides hints to the calling
  code that can be used to avoid creating reduction and reshaping ops.

  Args:
    x: The left input tensor to a broadcasting binary op.
    y: The right input tensor to a broadcasting binary op.
    grad: The incoming gradient tensor for a broadcasting binary op.

  Returns:
    A pair of tuples, containing:
      * A 3-tuple of broadcast information for x, containing:
        * The shape of x (as a tuple or Tensor).
        * The reduction indices for x (as a tuple or Tensor).
        * A boolean, which if True, indicates that x's shape differs from grad's
          shape (and so x's gradient must be reduced and/or reshaped).
      * A 3-tuple of broadcast information for y, containing the respective
        details for y.
  TNF)Úoptimize)r   Úexecuting_eagerlyÚ
isinstancer   ÚTensorr	   r$   r
   Úbroadcast_gradient_argsÚ_shape_tupleÚshape_internalr   Úget_default_graphÚ_bcast_grad_args_cacheÚKeyErrorÚtupler   Útry_evaluate_constant)r   r   r   ÚsxÚsyÚrxÚryÚx_shape_tupleÚy_shape_tupleÚgrad_shape_tupleÚx_needs_reductionÚy_needs_reductionÚgÚrx_valueÚry_valuer   r   r   ÚSmartBroadcastGradientArgs@   sP   
ÿ
ÿ
þ

ÿÿ
ÿõrC   r   c                 C   s   |   ¡ tu S r   )r0   Ú_empty_tuple)r   r   r   r   Ú	_IsScalarŠ   s   rE   ÚSumc                 C   sü  | j d  ¡ }|durÅt | j d ¡}|durÅt|ƒ}t |t |¡¡rst 	¡ rKt ¡ }| 
¡  |¡}|du rJtjdg| tjd}| 
¡  ||¡ ndg| }t ||¡}d|vrctj|tjd}nt | j d ¡}t ||¡dgS d|vrÅt 	¡ sÅt ¡ }t| d¡ƒ}z|j||f \}	}
W n% ty¶   dd„ }|t ||¡ƒ}	|t||	ƒƒ}
|	|
f|j||f< Y nw t ||	¡}t ||
¡dgS t | j d ¡}|  d¡söt |¡ t || j d ¡}	W d  ƒ n1 sëw   Y  t ||	¡}t ||¡dgS )	zGradient for Sum.r   Nr   ©Údtypeéÿÿÿÿc                 S   s4   t  | ¡rt  | ¡}|d usJ ‚t|ƒS | }t|ƒS r   )r   Ú
is_tf_typer6   r5   )ÚtÚvaluer   r   r   ÚEvaluateAsTuple¸   s   

ÿz!_SumGrad.<locals>.EvaluateAsTupler    ) r%   r0   r   Úconstant_valueÚlenÚnpÚarray_equalÚaranger   r,   Úones_rank_cacheÚgetr   Úconstantr   Úint32Úputr	   r&   r$   Útiler   r2   r5   Ú_reduced_shape_cacher4   r   r#   r   r"   Úcolocate_withÚbroadcast_to)r   r   Úinput_0_shapeÚaxesÚrankÚctxÚ	new_shapeÚinput_shapeÚgraphr)   Útile_scalingrM   r   r   r   Ú_SumGradŽ   s`   €
ÿ
ÿÿÿñ
ÿýrd   c                 C   s¤   t  | jd ¡}| jd }|  d¡s(t || jd ¡}t  ||¡}t  ||¡}nt  |¡}t t 	|| jd ¡|j
¡}t  t || jd ¡|¡}t ||¡| dgS )z@Gradient for Min or Max. Amazingly it's precisely the same code.r   r    r   N)r	   r$   r%   r!   r"   r   r#   r&   ÚcastÚequalrH   Ú
reduce_sumÚdivide)r   r   ra   r   r)   Ú
indicatorsÚnum_selectedr   r   r   Ú_MinOrMaxGradÖ   s   


ÿrk   ÚMaxc                 C   ó
   t | |ƒS )zGradient for Max.©rk   r   r   r   r   Ú_MaxGradë   ó   
ro   ÚMinc                 C   rm   r   rn   r   r   r   r   Ú_MinGradñ   s   
rr   ÚMeanc           
      C   sÖ   t | |ƒd }| jd  ¡ }| jd  ¡ }|dur?|dur?d|vr?d|vr?t |¡}t |¡}|t|dƒ }tj||j	d}nt
 | jd ¡}t
 |¡}| jd | | }	t t
 ||	¡¡}t |t ||j	¡¡dfS )zGradient for Mean.r   Nr   rG   )rd   r%   r0   r!   rP   ÚprodÚmaxr   rU   rH   r	   r$   Úsizer   Úreduce_prodÚgatherr'   re   )
r   r   Úsum_gradra   Úoutput_shapeÚ
input_sizeÚoutput_sizeÚfactorÚ
input_rankr]   r   r   r   Ú	_MeanGradö   s   


r   ÚProdc                 C   s  t  | jd ¡}t  | jd dg¡}|  d¡s&t || jd ¡}t  ||¡}t  ||¡}t 	d¡G t  
| jd ¡}|| | }t |tj¡}t d|¡}t ||tj¡\}}	t  ||gd¡}
t t  ||¡¡}t t  ||¡¡}W d  ƒ n1 s{w   Y  t  | jd |
¡}t  |¡}t  |||f¡}tj|ddd}tj|dddd	}t  t |¡t |¡ |¡}|t  |t  |
¡¡ }t  ||¡dfS )
zGradient for Prod.r   r   rI   r    z/cpu:0NT)ÚaxisÚ	exclusive)r   r‚   Úreverse)r	   r$   r%   r&   r"   r   r#   r[   r   Údevicer^   re   r   rV   Úranger
   Ú	list_diffÚconcatrw   rx   Ú	transposeÚcumprodÚconjÚinvert_permutation)r   r   ra   Úreduction_indicesr)   r^   ÚreducedÚidxÚotherÚ_ÚpermÚreduced_numÚ	other_numÚpermutedÚpermuted_shapeÚreshapedÚleftÚrightr   Úoutr   r   r   Ú	_ProdGrad
  s4   
ø	
ÿrš   Ú
SegmentSumc                 C   s   t  || jd ¡dfS )zGradient for SegmentSum.r   N)r	   rx   r%   r   r   r   r   Ú_SegmentSumGrad;  s   rœ   ÚSegmentMeanc                 C   s„   t  | jd ¡}t  t  | jd ¡t jt  |d d¡tjdgd¡}t j||j	d}t
 |t
 || jd ¡¡}t  || jd ¡dfS )zGradient for SegmentMean.r   r   rG   N)r	   r^   r%   r‡   r$   ÚonesÚexpand_dimsr   rV   rH   r   rh   Úsegment_sumrx   )r   r   r~   Ú
ones_shaperž   Úscaled_gradr   r   r   Ú_SegmentMeanGradA  s   ÿþür£   ÚSparseSegmentSumc                 C   sl   t  | jd ¡d }t ddd¡r"t || jd | jd |¡ddfS t t  || jd ¡| jd |¡ddfS )zGradient for SparseSegmentSum.r   éå  é   é
   r   é   N©	r	   r$   r%   r   Úforward_compatibler   Úsparse_segment_sum_gradÚunsorted_segment_sumrx   ©r   r   Údim0r   r   r   Ú_SparseSegmentSumGradO  s   ÿÿÿÿr¯   ÚSparseSegmentSumWithNumSegmentsc                 C   sp   t  | jd ¡d }t ddd¡r#t || jd | jd |¡dddfS t t  || jd ¡| jd |¡dddfS )z-Gradient for SparseSegmentSumWithNumSegments.r   r¥   r¦   r§   r   r¨   Nr©   r­   r   r   r   Ú$_SparseSegmentSumWithNumSegmentsGrad[  s   ÿÿþþr±   ÚSparseSegmentMeanc                 C   ó6   t  | jd ¡d }t || jd | jd |¡ddfS )zGradient for SparseSegmentMean.r   r   r¨   N©r	   r$   r%   r   Úsparse_segment_mean_gradr­   r   r   r   Ú_SparseSegmentMeanGradh  ó   ÿÿr¶   Ú SparseSegmentMeanWithNumSegmentsc                 C   ó8   t  | jd ¡d }t || jd | jd |¡dddfS )z.Gradient for SparseSegmentMeanWithNumSegments.r   r   r¨   Nr´   r­   r   r   r   Ú%_SparseSegmentMeanWithNumSegmentsGradp  ó   ÿÿrº   ÚSparseSegmentSqrtNc                 C   r³   )z Gradient for SparseSegmentSqrtN.r   r   r¨   N©r	   r$   r%   r   Úsparse_segment_sqrt_n_gradr­   r   r   r   Ú_SparseSegmentSqrtNGradx  r·   r¿   Ú!SparseSegmentSqrtNWithNumSegmentsc                 C   r¹   )z/Gradient for SparseSegmentSqrtNWithNumSegments.r   r   r¨   Nr½   r­   r   r   r   Ú&_SparseSegmentSqrtNWithNumSegmentsGrad€  r»   rÁ   c                 C   s’   t j| jd | jd jd}t  | jd | jd ¡}t | jd |¡}t t 	||j¡| jd ¡}t 
||¡}t  || jd ¡}t  |||¡dfS )z) Gradient for SegmentMin and SegmentMax. r   rG   r   N)r	   Ú
zeros_liker%   rH   rx   r!   r   rf   r    re   rh   Úwhere_v2)r   r   ÚzerosÚgathered_outputsÚis_selectedrj   Úweighted_gradsÚgathered_gradsr   r   r   Ú_SegmentMinOrMaxGradˆ  s   ÿrÉ   Ú
SegmentMinc                 C   rm   )zGradient for SegmentMin.©rÉ   r   r   r   r   Ú_SegmentMinGrad—  rp   rÌ   Ú
SegmentMaxc                 C   rm   )zGradient for SegmentMax.rË   r   r   r   r   Ú_SegmentMaxGrad  rp   rÎ   ÚSegmentProdc                 C   sÀ   | j d }| j d }t |d¡}t tj|tjd|¡}t 	t 
|d¡t |¡|¡}t 	|t |¡|¡}t ||¡}t | jd |¡}t ||¡}	|| }
t 	||	|
¡}t ||¡}|| dfS )aé  Gradient for SegmentProd.

  The gradient can be expressed for each segment by dividing the segment's
  product by each element of the segment input tensor, but this approach can't
  deal with zeros in the input.
  Unlike reduce_prod we can't use cumsum here as individual segments may have
  a different number of elements. Therefore we consider three cases:
  1) A segment input contains no zeros and we can safely divide by the input
     tensor.
  2) A segment contains exactly one zero. Then the gradient of each input of
     the segment is zero except for the 0-input, there the gradient is
     the product of the remaining segment entries.
  3) A segment contains at least two zeros. The gradient is zero for all
     segment inputs.
  r   r   rG   N)r%   r   rf   r   r    re   r   rV   r	   rÃ   ÚgreaterrÂ   Ú	ones_likeÚsegment_prodrx   r!   )r   r   ÚdataÚsegment_idsÚis_zeroÚ	num_zerosÚnon_zero_dataÚnon_zero_prodÚgathered_prodÚgathered_non_zero_prodÚprod_divided_by_elÚpartial_derivativeÚgathered_gradr   r   r   Ú_SegmentProdGrad£  s&   

ÿÿÿrÞ   c                 C   s²   |du rt  |t |¡¡}t | |¡}|du rJt  |d¡}t |¡}tj|tjt 	|¡t 	|¡ g|j
dgdd}t ||¡}|tj|tjd@ }t |¡}t |||¡||fS )an   Helper function for unsorted segment ops.

  Gathers params for
      positive segment ids and gathers 0 for inputs with negative segment id.
      Also returns the clipped indices and a boolean mask with the same shape
      as ids where a positive id is masked as true. With this, the latter two
      can be passed as arguments to this function to reuse them.
  Nr   rG   )r   )r   r   r	   rÂ   rx   Úgreater_equalr$   r‡   rž   r^   rH   r&   rÑ   r   ÚboolrÃ   )ÚparamsÚidsÚzero_clipped_indicesÚis_positiveÚgatheredÚis_positive_shapeÚbroadcastable_shapeÚ
zero_slicer   r   r   Ú_GatherDropNegativesË  s2   
ÿþÿûÿ
ÿÿré   c                 C   sœ   t | jd | jd ƒ\}}}t | jd |¡}t ||¡}t t ||j¡| jd | jd ¡}t 	||¡}t |d||ƒ\}}	}	t
 |¡}
t
 |||
¡ddfS )z9 Gradient for UnsortedSegmentMin and UnsortedSegmentMax. r   r   r¨   N)ré   r!   r%   r   rf   Úlogical_andr¬   re   rH   rh   r	   rÂ   rÃ   )r   r   rÅ   rã   rä   rÆ   rj   rÇ   rÈ   r   rÄ   r   r   r   Ú_UnsortedSegmentMinOrMaxGradî  s   ÿÿ
ÿ
rë   ÚUnsortedSegmentSumc                 C   s   t || jd ƒd ddfS )z Gradient for UnsortedSegmentSum.r   r   N)ré   r%   r   r   r   r   Ú_UnsortedSegmentSumGrad   s   rí   ÚUnsortedSegmentMaxc                 C   rm   )z" Gradient for UnsortedSegmentMax. ©rë   r   r   r   r   Ú_UnsortedSegmentMaxGrad  rp   rð   ÚUnsortedSegmentMinc                 C   rm   )z" Gradient for UnsortedSegmentMin. rï   r   r   r   r   Ú_UnsortedSegmentMinGrad  rp   rò   ÚUnsortedSegmentProdc                 C   s
  t  | jd d¡}t t j|tjd| jd | jd ¡}t 	t  
|d¡t |¡|¡}t 	|t | jd ¡| jd ¡}t || jd | jd ¡}t  | jd t | jd ¡¡}t | jd |¡}t ||¡}|| jd  }	t 	|||	¡}
t|| jd |ƒd }||
 ddfS )aò   Gradient for UnsortedSegmentProd.

  The gradient can be expressed for each segment by dividing the segment's
  product by each element of the segment input tensor, but this approach can't
  deal with zeros in the input.
  Unlike reduce_prod we can't use cumsum here as individual segments may have
  a different number of elements. Therefore we consider three cases:
  1) A segment input contains no zeros and we can safely divide by the input
     tensor.
  2) A segment contains exactly one zero. Then the gradient of each input of
     the segment is zero except for the 0-input, there the gradient is
     the product of the remaining segment entries.
  3) A segment contains at least two zeros. The gradient is zero for all
     segment inputs.
  r   rG   r   r¨   N)r   rf   r%   r   r¬   re   r   rV   r	   rÃ   rÐ   rÂ   rÑ   Úunsorted_segment_prodr   rx   r!   ré   )r   r   rÕ   rÖ   r×   rØ   rã   rÙ   rÚ   rÛ   rÜ   rÝ   r   r   r   Ú_UnsortedSegmentProdGrad  s8   ÿÿÿÿÿÿÿÿrõ   ÚAbsc                 C   s   | j d }|t |¡ S ©Nr   )r%   r   Úsign©r   r   r   r   r   r   Ú_AbsGradB  s   
rú   ÚNegc                 C   s   | S )zReturns -grad.r   ©r   r   r   r   r   Ú_NegGradH  ó   rý   ÚInvc                 C   ó   | j d }t ||¡S ©zReturns -grad * (1 / x^2).r   ©r!   r   Úreciprocal_grad©r   r   r   r   r   r   Ú_InvGradN  ó   
r  Ú
Reciprocalc                 C   r   r  r  r  r   r   r   Ú_ReciprocalGradU  r  r  ÚInvGradc                 C   óp   | j d }t |g¡# t | j d ¡}t |¡}|d | | t ||¡fW  d   ƒ S 1 s1w   Y  d S ©Nr   r   ç       À©r%   r   Úcontrol_dependenciesr   rŠ   r   r  ©r   r   ÚbÚcaÚcgr   r   r   Ú_InvGradGrad\  ó   

$ýr  ÚReciprocalGradc                 C   r
  r  r  r  r   r   r   Ú_ReciprocalGradGradf  r  r  ÚSquarec                 C   sh   | j d }t |g¡ t |¡}tjd|jd}t |t ||¡¡W  d   ƒ S 1 s-w   Y  d S )Nr   ç       @rG   )	r%   r   r  r   rŠ   r   rU   rH   Úmultiply©r   r   r   r   r   r   r   Ú_SquareGradp  s   

$ýr  ÚSqrtc                 C   r   r÷   )r!   r   Ú	sqrt_gradr  r   r   r   Ú	_SqrtGradz  s   
r  ÚSqrtGradc                 C   sd   | j d }| jd }t |g¡ || }t |¡ | d| fW  d   ƒ S 1 s+w   Y  d S )Nr   ç      à?)r%   r!   r   r  r   rŠ   )r   r   Úar   Úgar   r   r   Ú_SqrtGradGrad€  s   

$þr#  ÚRsqrtc                 C   r   )z Returns -0.5 * grad * conj(y)^3.r   )r!   r   Ú
rsqrt_gradr  r   r   r   Ú
_RsqrtGrad‰  r  r&  Ú	RsqrtGradc                 C   s‚   | j d }| j d }t |g¡' t |¡}t |¡}d| | t |¡ }t ||¡}||fW  d  ƒ S 1 s:w   Y  dS )z<Returns backprop gradient for f(a,b) = -0.5 * b * conj(a)^3.r   r   g      ø¿N)r%   r   r  r   rŠ   Úsquarer   r%  )r   r   r!  r  r  r  Úgrad_aÚgrad_br   r   r   Ú_RsqrtGradGrad  s   



$ûr+  ÚExpc                 C   sL   | j d }t |g¡ t |¡}|| W  d  ƒ S 1 sw   Y  dS ©zReturns grad * exp(x).r   N)r!   r   r  r   rŠ   r  r   r   r   Ú_ExpGrad  ó
   

$þr.  ÚExpm1c                 C   sV   | j d }t |g¡ t |¡}t |¡}|| W  d  ƒ S 1 s$w   Y  dS r-  )r%   r   r  r   rŠ   Úexpr  r   r   r   Ú
_Expm1Grad¦  s   


$ýr2  ÚLogc                 C   óR   | j d }t |g¡ t |¡}|t |¡ W  d  ƒ S 1 s"w   Y  dS )zReturns grad * (1/x).r   N©r%   r   r  r   rŠ   Ú
reciprocalrù   r   r   r   Ú_LogGrad°  ó
   

$þr7  ÚLog1pc                 C   sV   | j d }t |g¡ t |¡}|t d| ¡ W  d  ƒ S 1 s$w   Y  dS )zReturns grad * (1/(1 + x)).r   r   Nr5  rù   r   r   r   Ú
_Log1pGrad¹  s
   

$þr:  ÚXlogyc              	   C   sÔ   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡> tjt 	|tjd|j
d¡|j
d}t ||¡}	t ||¡}
t t |	| |¡|¡t t |
| |¡|¡fW  d  ƒ S 1 scw   Y  dS )z8Returns gradient of xlogy(x, y) with respect to x and y.r   r   ç        rG   N)r%   r	   r$   r
   r/   r   r  r   re   Ú	not_equalrH   r   ÚxlogyÚxdivyr&   rg   ©r   r   r   r   r7   r8   r9   r:   Ú
not_zero_xÚ	partial_xÚ	partial_yr   r   r   Ú
_XLogyGradÂ  s   



ÿÿ$ûrD  ÚXlog1pyc              	   C   sØ   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡@ tjt 	|tjd|j
d¡|j
d}t ||¡}	t ||d ¡}
t t |	| |¡|¡t t |
| |¡|¡fW  d  ƒ S 1 sew   Y  dS )z:Returns gradient of xlog1py(x, y) with respect to x and y.r   r   r<  rG   ç      ð?N)r%   r	   r$   r
   r/   r   r  r   re   r=  rH   r   Úxlog1pyr?  r&   rg   r@  r   r   r   Ú_XLog1pyGradÓ  s   



ÿÿ$ûrH  ÚXdivyc              	   C   sÞ   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡C tjt 	|tjd|j
d¡|j
d}t ||¡}	t t |¡|d ¡}
t t |	| |¡|¡t t |
| |¡|¡fW  d  ƒ S 1 shw   Y  dS )z8Returns gradient of xdivy(x, y) with respect to x and y.r   r   r<  rG   r¨   N)r%   r	   r$   r
   r/   r   r  r   re   r=  rH   r   r?  Únegativer&   rg   r@  r   r   r   Ú
_XDivyGradä  s   



ÿÿ$ûrK  ÚSinhc                 C   r4  )zReturns grad * cosh(x).r   N)r%   r   r  r   rŠ   Úcoshrù   r   r   r   Ú	_SinhGradõ  r8  rN  ÚCoshc                 C   r4  )zReturns grad * sinh(x).r   N)r%   r   r  r   rŠ   Úsinhrù   r   r   r   Ú	_CoshGradþ  r8  rQ  ÚTanhc                 C   óP   | j d }t |g¡ t |¡}t ||¡W  d  ƒ S 1 s!w   Y  dS )z'Returns grad * (1 - tanh(x) * tanh(x)).r   N)r!   r   r  r   rŠ   r   Ú	tanh_gradr  r   r   r   Ú	_TanhGrad  ó
   


$þrU  ÚAsinhc                 C   óR   | j d }t |g¡ t |¡}|t |¡ W  d  ƒ S 1 s"w   Y  dS )zReturns grad * 1/cosh(y).r   N)r!   r   r  r   rŠ   rM  r  r   r   r   Ú
_AsinhGrad  r8  rY  ÚAcoshc                 C   rX  )zReturns grad * 1/sinh(y).r   N)r!   r   r  r   rŠ   rP  r  r   r   r   Ú
_AcoshGrad  r8  r[  ÚAtanhc                 C   óx   | j d }t |g¡' t |¡}t |¡}tjd|jd}t 	t 
||¡¡}|| W  d  ƒ S 1 s5w   Y  dS )zReturns grad * 1/ (1 - x^2).r   r   rG   N)r%   r   r  r   rŠ   r(  r   rU   rH   r6  Úsubtract©r   r   r   Úx2ÚoneÚinvr   r   r   Ú
_AtanhGrad"  ó   


$ûrc  ÚTanhGradc                 C   sl   t  |g¡& t | jd ¡}t | jd ¡}|d | | t ||¡fW  d   ƒ S 1 s/w   Y  d S )Nr   r   r  )r   r  r   rŠ   r%   r   rT  )r   r   r!  r  r   r   r   Ú_TanhGradGrad.  s
   $ýrf  ÚErfc                 C   óz   | j d }tjdt tj¡ |jd}t |g¡ t	 
|¡}|| t	 t	 |¡ ¡ W  d  ƒ S 1 s6w   Y  dS )z'Returns grad * 2/sqrt(pi) * exp(-x**2).r   r¨   rG   N©r%   r   rU   rP   ÚsqrtÚpirH   r   r  r   rŠ   r1  r(  )r   r   r   Útwo_over_root_pir   r   r   Ú_ErfGrad6  s   

$þrm  ÚErfcc                 C   rh  )z(Returns -grad * 2/sqrt(pi) * exp(-x**2).r   éþÿÿÿrG   Nri  )r   r   r   Úminus_two_over_root_pir   r   r   Ú	_ErfcGrad@  s   
ÿ
$þrq  ÚErfinvc                 C   sj   t jt tj¡d |jd}t |g¡ || t 	t 
| jd ¡¡ W  d  ƒ S 1 s.w   Y  dS )z0Returns grad * sqrt(pi) / 2 * exp(erfinv(x)**2).r¨   rG   r   N©r   rU   rP   rj  rk  rH   r   r  r   r1  r(  r!   )r   r   Úroot_pi_over_twor   r   r   Ú_ErfinvGradK  s   
ÿ$ÿru  ÚNdtric                 C   sn   t jt dtj ¡|jd}t |g¡ || t 	t 
| jd ¡d ¡ W  d  ƒ S 1 s0w   Y  dS )z3Returns grad * sqrt(2 * pi) * exp(ndtri(x)**2 / 2).r¨   rG   r   r  Nrs  )r   r   Úroot_two_pir   r   r   Ú
_NdtriGradT  s   
ÿ$ÿrx  ÚLgammac                 C   r4  )zReturns grad * digamma(x).r   N)r%   r   r  r   rŠ   Údigammarù   r   r   r   Ú_LgammaGrad]  r8  r{  ÚDigammac                 C   sd   | j d }t |g¡ t |¡}t tjd|jd|¡}|| W  d  ƒ S 1 s+w   Y  dS )zFCompute gradient of the digamma function with respect to its argument.r   r   rG   N)	r%   r   r  r   rŠ   Ú	polygammar	   rU   rH   ©r   r   r   rB  r   r   r   Ú_DigammaGradf  s   

$ýr  ÚDawsnc                 C   sX   | j d }| jd }t |g¡ |dd| |   W  d  ƒ S 1 s%w   Y  dS )z:Compute gradient of dawsn(x) with respect to its argument.r   rF  r¨   N)r%   r!   r   r  r  r   r   r   Ú
_DawsnGradp  s
   

$ÿr  ÚExpintc                 C   sL   | j d }t |g¡ |t |¡ | W  d  ƒ S 1 sw   Y  dS )z;Compute gradient of expint(x) with respect to its argument.r   N)r%   r   r  r   r1  rù   r   r   r   Ú_ExpintGrady  s   
$ÿrƒ  Ú
FresnelCosc                 C   óX   | j d }t |g¡ |t tjd t |¡ ¡ W  d  ƒ S 1 s%w   Y  dS )z@Compute gradient of fresnel_cos(x) with respect to its argument.r   r  N)r%   r   r  r   ÚcosrP   rk  r(  rù   r   r   r   Ú_FresnelCosGrad  ó   
$ÿr‡  Ú
FresnelSinc                 C   r…  )z@Compute gradient of fresnel_sin(x) with respect to its argument.r   r  N)r%   r   r  r   ÚsinrP   rk  r(  rù   r   r   r   Ú_FresnelSinGrad‰  rˆ  r‹  ÚSpencec                 C   sr   | j d }t |g¡$ t |¡d|  }t t |d¡t |¡ |¡}|| W  d  ƒ S 1 s2w   Y  dS )z;Compute gradient of spence(x) with respect to its argument.r   r   rF  N)	r%   r   r  r   Úlogr	   Úwhererf   rÑ   r~  r   r   r   Ú_SpenceGrad‘  s   
ÿ$ür  ÚBesselI0c                 C   sL   | j d }t |g¡ t |¡}|| W  d  ƒ S 1 sw   Y  dS )z>Compute gradient of bessel_i0(x) with respect to its argument.r   N)r%   r   r  r   Ú	bessel_i1r~  r   r   r   Ú_BesselI0Gradœ  r/  r’  Ú	BesselI0ec                 C   sd   | j d }| jd }t |g¡ t |¡t |¡|  }|| W  d  ƒ S 1 s+w   Y  dS )z?Compute gradient of bessel_i0e(x) with respect to its argument.r   N)r%   r!   r   r  r   Ú
bessel_i1er   rø   ©r   r   r   r   rB  r   r   r   Ú_BesselI0eGrad¥  s   

$þr–  ÚBesselI1c              
   C   ó~   | j d }| jd }t |g¡% t t |d¡t d|j	¡t
 |¡t ||¡ ¡}|| W  d  ƒ S 1 s8w   Y  dS )z>Compute gradient of bessel_i1(x) with respect to its argument.r   r<  rF  N)r%   r!   r   r  r	   rÃ   r   rf   re   rH   r   Ú	bessel_i0Údiv©r   r   r   r   Údy_dxr   r   r   Ú_BesselI1Grad¯  ó   

þ$÷r  Ú	BesselI1ec                 C   sŠ   | j d }| jd }t |g¡+ t t |d¡t d|j	¡t
 |¡|t |¡t |¡   ¡}|| W  d  ƒ S 1 s>w   Y  dS )z?Compute gradient of bessel_i1e(x) with respect to its argument.r   r<  r   N)r%   r!   r   r  r	   rÃ   r   rf   re   rH   r   Ú
bessel_i0erø   r6  r›  r   r   r   Ú_BesselI1eGradÀ  s   


ÿþ$ör¡  ÚBesselK0c                 C   óN   | j d }t |g¡ t |¡ }|| W  d  ƒ S 1 s w   Y  dS )z>Compute gradient of bessel_k0(x) with respect to its argument.r   N)r%   r   r  r   Ú	bessel_k1r~  r   r   r   Ú_BesselK0GradÒ  ó
   
$þr¥  Ú	BesselK0ec                 C   sZ   | j d }| jd }t |g¡ |t |¡ }|| W  d  ƒ S 1 s&w   Y  dS )z?Compute gradient of bessel_k0e(x) with respect to its argument.r   N)r%   r!   r   r  r   Ú
bessel_k1er•  r   r   r   Ú_BesselK0eGradÛ  s   

$þr©  ÚBesselK1c                 C   sd   | j d }| jd }t |g¡ t |¡ t ||¡ }|| W  d  ƒ S 1 s+w   Y  dS )z>Compute gradient of bessel_k1(x) with respect to its argument.r   N)r%   r!   r   r  r   Ú	bessel_k0r   rš  r•  r   r   r   Ú_BesselK1Gradå  s   

$ür¬  Ú	BesselK1ec                 C   sh   | j d }| jd }t |g¡ |dt |¡  t |¡ }|| W  d  ƒ S 1 s-w   Y  dS )z?Compute gradient of bessel_k1e(x) with respect to its argument.r   rF  N)r%   r!   r   r  r   r6  r   Ú
bessel_k0er•  r   r   r   Ú_BesselK1eGradñ  s   

ÿ$ûr¯  ÚBesselJ0c                 C   r£  )z>Compute gradient of bessel_j0(x) with respect to its argument.r   N)r%   r   r  r   Ú	bessel_j1r~  r   r   r   Ú_BesselJ0Gradþ  r¦  r²  ÚBesselJ1c              
   C   r˜  )z>Compute gradient of bessel_j1(x) with respect to its argument.r   r<  r   N)r%   r!   r   r  r	   rÃ   r   rf   re   rH   r   Ú	bessel_j0rš  r›  r   r   r   Ú_BesselJ1Grad  rž  rµ  ÚBesselY0c                 C   r£  )z>Compute gradient of bessel_y0(x) with respect to its argument.r   N)r%   r   r  r   Ú	bessel_y1r~  r   r   r   Ú_BesselY0Grad  r¦  r¸  ÚBesselY1c                 C   sb   | j d }| jd }t |g¡ t |¡t ||¡ }|| W  d  ƒ S 1 s*w   Y  dS )z>Compute gradient of bessel_y1(x) with respect to its argument.r   N)r%   r!   r   r  r   Ú	bessel_y0r   rš  r•  r   r   r   Ú_BesselY1Grad!  s   

$ür»  ÚIgammac           
      C   sÌ   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡: t ||¡}t	 
| |d t	 |¡  t	 |¡ ¡}	t t	 || |¡|¡t t	 |	| |¡|¡fW  d  ƒ S 1 s_w   Y  dS )z9Returns gradient of igamma(a, x) with respect to a and x.r   r   N)r%   r	   r$   r
   r/   r   r  r   Úigamma_grad_ar   r1  r  Úlgammar&   rg   )
r   r   r!  r   Úsar7   Úrar9   Ú	partial_arB  r   r   r   Ú_IgammaGrad-  s   



ÿÿ$úrÂ  ÚIgammacc                 C   s   t | |ƒ\}}| | fS )zDReturns gradient of igammac(a, x) = 1 - igamma(a, x) w.r.t. a and x.)rÂ  )r   r   r½  Úigamma_grad_xr   r   r   Ú_IgammacGrad@  s   rÅ  ÚBetaincc                 C   sœ   | j \}}}t |¡}t |¡}t ||¡\}}t |¡t |¡ t || ¡ }	t t 	|d | ¡t 
|d |¡ |	 ¡}
ddt t |
| |¡|¡fS )z7Returns gradient of betainc(a, b, x) with respect to x.r   N)r%   r	   r$   r
   r/   r   r¾  r   r1  rG  r>  r&   rg   )r   r   r!  r  r   r¿  r7   r   r9   Úlog_betarB  r   r   r   Ú_BetaincGradG  s"   

ÿÿÿÿýrÈ  ÚZetac           	      C   s®   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡+ t |¡}t |¡}| t 	|d |¡ }dt 
t || |¡|¡fW  d  ƒ S 1 sPw   Y  dS )z7Returns gradient of zeta(x, q) with respect to x and q.r   r   N)r%   r	   r$   r
   r/   r   r  r   rŠ   Úzetar&   rg   )	r   r   r   Úqr7   ÚsqÚ	unused_rxÚrqÚ	partial_qr   r   r   Ú	_ZetaGradc  s   





ÿ$ürÐ  Ú	Polygammac           	      C   s¨   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡( t |¡}t |¡}t 	|d |¡}dt 
t || |¡|¡fW  d  ƒ S 1 sMw   Y  dS )z6Returns gradient of psi(n, x) with respect to n and x.r   r   N)r%   r	   r$   r
   r/   r   r  r   rŠ   r}  r&   rg   )	r   r   Únr   Úsnr7   Ú	unused_rnr9   rB  r   r   r   Ú_PolygammaGradv  s   





ÿ$ürÕ  ÚSigmoidc                 C   rS  )z-Returns grad * sigmoid(x) * (1 - sigmoid(x)).r   N)r!   r   r  r   rŠ   r   Úsigmoid_gradr  r   r   r   Ú_SigmoidGrad‰  rV  rØ  ÚSigmoidGradc                 C   st   t  |g¡* t | jd ¡}t | jd ¡}|| }|d| |  t ||¡fW  d   ƒ S 1 s3w   Y  d S )Nr   r   r  )r   r  r   rŠ   r%   r   r×  )r   r   r!  r  Úgbr   r   r   Ú_SigmoidGradGrad’  s   $ürÛ  ÚSignc                 C   s   | j d }t |¡S )z
Returns 0.r   )r%   r	   rÂ   )r   r   r   r   r   r   Ú	_SignGrad›  s   

rÝ  ÚSinc                 C   r4  )zReturns grad * cos(x).r   N)r%   r   r  r   rŠ   r†  rù   r   r   r   Ú_SinGrad¢  r8  rß  ÚCosc                 C   sT   | j d }t |g¡ t |¡}| t |¡ W  d  ƒ S 1 s#w   Y  dS )zReturns grad * -sin(x).r   N)r%   r   r  r   rŠ   rŠ  rù   r   r   r   Ú_CosGrad«  s
   

$þrá  ÚTanc                 C   sf   | j d }t |g¡ t |¡}t t |¡¡}t |¡}|| W  d  ƒ S 1 s,w   Y  dS )zReturns grad * 1/sec^2(x).r   N)r%   r   r  r   rŠ   r6  r†  r(  )r   r   r   ÚsecxÚsecx2r   r   r   Ú_TanGrad´  s   


$ürå  ÚAsinc                 C   s‚   | j d }t |g¡, t |¡}t |¡}tjd|jd}t 	t 
||¡¡}t |¡}|| W  d  ƒ S 1 s:w   Y  dS )zReturns grad * 1/sqrt(1-x^2).r   r   rG   N©r%   r   r  r   rŠ   r(  r   rU   rH   rj  r^  r6  ©r   r   r   r`  ra  Údenrb  r   r   r   Ú	_AsinGrad¿  s   



$úrê  ÚAcosc                 C   s„   | j d }t |g¡- t |¡}t |¡}tjd|jd}t 	t 
||¡¡}t |¡}| | W  d  ƒ S 1 s;w   Y  dS )zReturns grad * -1/sqrt(1-x^2).r   r   rG   Nrç  rè  r   r   r   Ú	_AcosGradÌ  s   



$úrì  ÚAtanc                 C   r]  )zReturns grad * 1/ (1 + x^2).r   r   rG   N)r%   r   r  r   rŠ   r(  r   rU   rH   r6  Úaddr_  r   r   r   Ú	_AtanGradÙ  rd  rï  ÚAtan2c                 C   sÂ   | j d }| j d }t |g¡G t|||ƒ\\}}}\}}}	|t |¡t |¡  }
| |
 }|r<t t ||¡|¡}||
 }|	rLt t ||¡|¡}||fW  d  ƒ S 1 sZw   Y  dS )z8Returns grad * x / (x^2 + y^2), grad * -y / (x^2 + y^2).r   r   N)	r%   r   r  rC   r   r(  r	   r&   rg   )r   r   r   r   r7   r9   Úmust_reduce_xr8   r:   Úmust_reduce_yÚgrad_invÚgxÚgyr   r   r   Ú
_Atan2Gradå  s   


ÿ
$òrö  ÚAddNc                 C   s   |gt | jƒ S )z"Copies the gradient to all inputs.)rO   r%   r   r   r   r   Ú	_AddNGradû  s   rø  c                 C   s8   |   ¡ }|  ¡ }|  ¡ }||ko||ko|d uod |vS r   )r0   )r   r   r   Úx_shapeÚy_shapeÚ
grad_shaper   r   r   Ú_ShapesFullySpecifiedAndEqual  s   ÿÿrü  ÚAddÚAddV2c                 C   s  | j d }d}z| j}|durd|v rt|ƒr|dfW S W n	 ty&   Y nw | j d }t|tjƒr<t|||ƒr<||fS t|||ƒ\\}}}\}}	}
|durUd|v rUd}n|sZ|}n
t	 
t ||¡|¡}|durrd|v rrd}||fS |
sz|}||fS t	 
t ||	¡|¡}||fS )zGradient for Add.r   Nr   ©r%   Úskip_input_indicesrE   ÚAttributeErrorr-   r   r.   rü  rC   r	   r&   r   rg   ©r   r   r   r   r   r7   r9   rñ  r8   r:   rò  rô  rõ  r   r   r   Ú_AddGrad  s@   
ÿ
€þ

ÿ
ÿüÿr  ÚSubc                 C   s  | j d }d}z| j}|durd|v rt|ƒr|dfW S W n	 ty&   Y nw | j d }t|tjƒr=t|||ƒr=|| fS t|||ƒ\\}}}\}}	}
|durVd|v rVd}n|s[|}n
t	 
t ||¡|¡}|dursd|v rsd}||fS |
s|| }||fS t	 
t | |	¡|¡}||fS )zGradient for Sub.r   Nr   rÿ  r  r   r   r   Ú_SubGrad/  s@   
ÿ
€þ

ÿ

ÿüÿr  ÚMulc                 C   s–  | j d }d}z| j}|dur#d|v r#t|ƒr#t |t |¡¡dfW S W n	 ty-   Y nw | j d }t|t	j
ƒrTt|||ƒrT|jtjtjfv rTt ||¡t ||¡fS |jj|jjkseJ |jd|jfƒ‚t|||ƒ\\}}}\}}	}
t |¡}t |¡}|durˆd|v rˆd}n|s‘t ||¡}nt t t ||¡|¡|¡}|dur­d|v r­d}||fS |
s¹t ||¡}||fS t t t ||¡|	¡|¡}||fS )z&The gradient of scalar multiplication.r   Nr   ú vs. )r%   r   rE   r   Úmulr   rŠ   r  r-   r   r.   rü  rH   r   rV   Úfloat32Ú
base_dtyperC   r	   r&   rg   r  r   r   r   Ú_MulGradQ  sP   
ÿ€þ

ÿ"
ÿ

ÿûþÿr  ÚMulNoNanc              	   C   sÂ   | j d }| j d }t|tjƒr"t|||ƒr"t ||¡t ||¡fS |jj|jjks3J |jd|jfƒ‚t	 
|¡}t	 
|¡}t ||¡\}}t	 t t ||¡|¡|¡t	 t t ||¡|¡|¡fS )z;The gradient of scalar multiplication with NaN-suppression.r   r   r  )r%   r-   r   r.   rü  r   Ú
mul_no_nanrH   r
  r	   r$   r
   r/   r&   r   rg   ©r   r   r   r   r7   r8   r9   r:   r   r   r   Ú_MulNoNanGradz  s"   


ÿ"

ÿÿþr  ÚDivc                 C   ó’   | j d }| j d }t |¡}t |¡}t ||¡\}}t |¡}t |¡}t t t 	||¡|¡|¡t t |t 	t 	| |¡|¡ |¡|¡fS )z"The gradient for the Div operator.r   r   )
r%   r	   r$   r
   r/   r   rŠ   r&   rg   rh   r  r   r   r   Ú_DivGradŒ  s   





ÿþþr  ÚFloorDivc                 C   ó   dS )z'The gradient for the FloorDiv operator.©NNr   ©r   Úunused_gradr   r   r   Ú_FloorDivGradž  s   r  ÚFloorModc                 C   sŠ   t  | jd ¡}t  | jd ¡}t |¡}t |¡}t ||¡\}}t  ||¡}t t  	||¡|¡}	t t  	|t  
|¡ |¡|¡}
|	|
fS )z Returns grad * (1, -floor(x/y)).r   r   )r   rŠ   r%   r	   r$   r
   r/   Ú	floor_divr&   rg   rJ  )r   r   r   r   r7   r8   r9   r:   Úfloor_xyrô  rõ  r   r   r   Ú_FloorModGrad¤  s   

ÿr  ÚTruncateDivc                 C   r  )Nr  r   r  r   r   r   Ú_TruncateDivGrad´  s   r  ÚRealDivc                 C   r  )zRealDiv op gradient.r   r   )
r%   r	   r$   r
   r/   r   rŠ   r&   rg   Úrealdivr  r   r   r   Ú_RealDivGrad¹  s"   





ÿÿþþr!  ÚDivNoNanc                 C   r  )zDivNoNan op gradient.r   r   )
r%   r	   r$   r
   r/   r   rŠ   r&   rg   Ú
div_no_nanr  r   r   r   Ú_DivNoNanGradÊ  s$   





ÿþüýr$  ÚPowc                 C   sš  | j d }| j d }d}z*| j}|dur5d|v r5t|ƒr5t |¡}t |¡}|| t ||d ¡ dfW S W n	 ty?   Y nw t|||ƒ\\}}}\}}	}
t |¡}t |¡}|du s`d|vry|| t ||d ¡ }|rxt 	t 
||¡|¡}nd}|du sƒd|vrÇt | jd ¡}|jjr–t |d¡}n|dk}t ||t |¡¡}t |t |¡t |¡¡}|| | }|
rÃt 	t 
||	¡|¡}||fS d}||fS )z%Returns grad * (y*x^(y-1), z*log(x)).r   r   N)r%   r   rE   r   rŠ   Úpowr  rC   r	   r&   rg   r!   rH   Ú
is_complexr=  rŽ  rÑ   r  rÂ   )r   r   r   r   r   r7   r9   rñ  r8   r:   rò  rô  ÚzÚmaskÚsafe_xÚlog_xrõ  r   r   r   Ú_PowGradÞ  sL   

ÿ

€þ
ÿ

€þr,  c           	      C   sB   | j d }| j d }t |¡}|||ƒ}t |||¡}d }||fS ©Nr   r   )r%   r	   rÂ   rÃ   )	r   r   Úselector_opr   r   rÄ   ÚxmaskÚxgradÚygradr   r   r   Ú_MaximumMinimumGradInputOnly  s   



r2  c                 C   s  | j d }d}z| j}|durd|v rt|ƒrt| ||ƒW S W n	 ty(   Y nw | j d }t |¡}t |¡}t |¡}|||ƒ}	t 	||¡\}
}|durUd|v rUd}nt 
|	||¡}t t ||
¡|¡}|durtd|v rtd}||fS t 
|	||¡}t t ||¡|¡}||fS )z;Factor out the code for the gradient of Maximum or Minimum.r   Nr   )r%   r   rE   r2  r  r	   r$   rÂ   r
   r/   rÃ   r&   r   rg   )r   r   r.  r   r   r   r7   r8   rÄ   r/  r9   r:   rô  r0  rõ  r1  r   r   r   Ú_MaximumMinimumGrad  s8   
ÿ€þ




ýr3  ÚMaximumc                 C   ó   t | |tjƒS )z/Returns grad*(x >= y, x < y) with type of grad.)r3  r   rß   r   r   r   r   Ú_MaximumGrad@  ó   r6  ÚMinimumc                 C   r5  )z/Returns grad*(x <= y, x > y) with type of grad.)r3  r   Ú
less_equalr   r   r   r   Ú_MinimumGradF  r7  r:  ÚSquaredDifferencec                 C   s4  | j d }| j d }d}z| j}W n	 ty   Y nw t |g¡ t d|¡||  }W d  ƒ n1 s6w   Y  t|tj	ƒrLt
|||ƒrL|| fS t|||ƒ\\}}}\}	}
}|dured|v red}n|rrt t ||¡|¡}n|}|dur‚d|v r‚d}||fS |r“t t ||
¡|	¡ }||fS | }||fS )z!Returns the gradient for (x-y)^2.r   r   Nr  )r%   r   r  r   r  r   Ú
scalar_mulr-   r   r.   rü  rC   r	   r&   rg   )r   r   r   r   r   Úx_gradr7   r9   rñ  r8   r:   rò  rô  rõ  r   r   r   Ú_SquaredDifferenceGradL  s<   


þý
ÿ

ÿüÿr>  ÚLessÚ	LessEqualÚGreaterÚGreaterEqualÚEqualÚApproximateEqualÚNotEqualÚ
LogicalAndÚ	LogicalOrÚ
LogicalNotÚSelectc                 C   s<   | j d }| j d }t |¡}d t |||¡t |||¡fS r-  )r%   r	   rÂ   rŽ  )r   r   Úcr   rÄ   r   r   r   Ú_SelectGrad  s   


ÿrK  ÚSelectV2c                 C   sÒ   | j d }| j d }| j d }tjg |jjd}t |||¡}t |¡}t | jd ¡}t 	||¡\}	}
t
j|d|	d}t ||¡}t |||¡}t |¡}t 	||¡\}}
t
j|d|d}t ||¡}d ||fS )Nr   r   r¨   rG   T)Úkeepdimsr   )r%   r	   rÄ   rH   r
  rÃ   r$   r!   r
   r/   r   rg   r&   )r   r   rJ  r   r   rÄ   rô  rù  rz   Úreduce_xr   rõ  rú  Úreduce_yr   r   r   Ú_SelectGradV2Š  s    





rP  c                 C   s¢   |   d¡}|   d¡}t | jd ¡}|s"|s"tj||dd}|dfS |s0|r0t ||¡}|dfS |r@|s@tj||dd}|dfS |rM|rMtj||ddd}|dfS )z.Gradient for MatMul, only for the first input.Útranspose_aÚtranspose_br   T©rR  ©rQ  rR  N©r"   r   rŠ   r%   r   Úmat_mul)r   r   Út_aÚt_br  r)  r   r   r   Ú_MatMulGradAgainstFirstOnly¢  s   

úüþrY  c                 C   s¢   |   d¡}|   d¡}t | jd ¡}|s"|s"tj||dd}d|fS |s2|r2tj||dd}d|fS |r@|s@t ||¡}d|fS |rM|rMtj||ddd}d|fS )z/Gradient for MatMul, only for the second input.rQ  rR  r   T©rQ  rT  NrU  )r   r   rW  rX  r!  r*  r   r   r   Ú_MatMulGradAgainstSecondOnly²  s   

úüþr[  ÚMatMulc           	      C   s>  z| j }|durd|v rt| |ƒW S d|v rt| |ƒW S W n	 ty&   Y nw |  d¡}|  d¡}t | jd ¡}t | jd ¡}|sY|sYtj	||dd}tj	||dd}||fS |so|rot 	||¡}tj	||dd}||fS |r…|s…tj	||dd}t 	||¡}||fS |r›|r›tj	||ddd	}tj	||ddd	}||fS )
zGradient for MatMul.Nr   r   rQ  rR  TrS  rZ  rT  )
r   rY  r[  r  r"   r   rŠ   r%   r   rV  )	r   r   r   rW  rX  r!  r  r)  r*  r   r   r   Ú_MatMulGradÂ  s>   €þ


÷úýr]  ÚSparseMatMulc                    s`  |   d¡}|   d¡}i ‰ |   d¡ˆ | jd  ¡ < |   d¡ˆ | jd  ¡ < t ¡  o.|jjdkˆ | ¡ < d‡ fd	d
„	}| jd j}| jd j}|s`|s`||| jd |dd|| jd ||ddfS |sx|rx||| jd |ƒ||| jd |ddfS |r|s|| jd ||dd|| jd ||ƒfS |r¬|r®|| jd ||ddd||| jd |dddfS dS dS )zGradient for SparseMatMul.rQ  rR  Úa_is_sparser   Úb_is_sparser   ÚReluGradFc                    sv   |   ¡ ˆ v r|  ¡ ˆ v sJ ‚ˆ |   ¡  }ˆ |  ¡  }|r#t |¡}d}tj| |||||d}|j|kr9t ||¡}|S )z*Helper function to create SparseMatMul op.F)rQ  rR  r_  r`  )Úrefr	   rˆ   r   ÚmatmulrH   re   )Út1Út2Ú	out_dtyperQ  rR  Ú	t1_sparseÚ	t2_sparsert   ©Ú	is_sparser   r   Ú_SparseMatMulð  s"   
ú
z(_SparseMatMulGrad.<locals>._SparseMatMulTrS  rZ  rT  N)FF)r"   r%   rb  r   r,   r   ÚtyperH   )r   r   rW  rX  rk  Údtype_aÚdtype_br   ri  r   Ú_SparseMatMulGradã  sB   




ÿÿÿÿÿþþÿro  ÚFloorc                 C   ó   d gS r   r   r  r   r   r   Ú
_FloorGrad  ó   rr  ÚCeilc                 C   rq  r   r   r  r   r   r   Ú	_CeilGrad  rs  ru  ÚRoundc                 C   rq  r   r   r  r   r   r   Ú
_RoundGrad!  rs  rw  ÚRintc                 C   rq  r   r   r  r   r   r   Ú	_RintGrad&  rþ   ry  ÚBatchMatMulc                 C   sä   | j d }| j d }|  d¡}|  d¡}|sD|s.tj||ddd}tj||ddd}||fS tj||ddd}tj||ddd}||fS |s\tj||ddd}tj||ddd}||fS tj||ddd}tj||ddd}||fS )ú<Returns the gradient of x and y given the gradient of x * y.r   r   Úadj_xÚadj_yFT©Ú	adjoint_aÚ	adjoint_b)r%   r"   r   rc  )r   r   r   r   r|  r}  Úgrad_xÚgrad_yr   r   r   Ú_BatchMatMul,  s&   



ö	ùýrƒ  ÚBatchMatMulV2ÚBatchMatMulV3c                 C   s®  | j d }| j d }|  d¡}|  d¡}|s>|s+tj||ddd}tj||ddd}n:tj||ddd}tj||ddd}n'|sStj||ddd}tj||ddd}ntj||ddd}tj||ddd}| ¡ }| ¡ }	|jdu p€|jd	kp€|	jdu p€|	jd	k}
|dd
…  ¡ oœ|	dd
…  ¡ oœ|dd
… |	dd
… k}|
r¡|r¥||fS t |¡}t |¡}t	 
|dd
… |dd
… ¡\}}t t ||¡|¡}t t ||¡|¡}||fS )r{  r   r   r|  r}  FTr~  Nr¨   ro  )r%   r"   r   rc  Ú	get_shaper^   Úis_fully_definedr	   r$   r
   r/   r&   rg   )r   r   r   r   r|  r}  r  r‚  Úshape_x_staticÚshape_y_staticÚ%output_may_have_non_empty_batch_shapeÚbatch_shapes_matchr7   r8   r9   r:   r   r   r   Ú_BatchMatMulV2F  sB   



þÿý

 rŒ  ÚRangeÚLinSpaceÚComplexc                 C   sl   | j d }| j d }t |¡}t |¡}t ||¡\}}t t t |¡|¡|¡t t t 	|¡|¡|¡fS )zBReturns the real and imaginary components of 'grad', respectively.r   r   )
r%   r	   r$   r
   r/   r&   r   rg   ÚrealÚimagr  r   r   r   Ú_ComplexGradx  s   



ÿr’  ÚRealc                 C   s   t jd|jd}t ||¡S )z=Returns 'grad' as the real part and set the imaginary part 0.r   rG   ©r   rU   rH   r   Úcomplex©r   r   Úzeror   r   r   Ú	_RealGrad„  ó   r˜  ÚImagc                 C   s   t jd|jd}t ||¡S )z=Returns 'grad' as the imaginary part and set the real part 0.r   rG   r”  r–  r   r   r   Ú	_ImagGrad‹  r™  r›  ÚAnglec                 C   s†   | j d }t |g¡. t |¡}t |¡}t t ||¡¡}tj	d|j
d}t ||¡}| | W  d  ƒ S 1 s<w   Y  dS )z Returns -grad / (Im(x) + iRe(x))r   rG   N)r%   r   r  r   r  r‘  r6  r•  r   rU   rH   )r   r   r   ÚreÚimr(  r—  Úcomplex_gradr   r   r   Ú
_AngleGrad’  s   


$úr   ÚConjc                 C   s
   t  |¡S )z&Returns the complex conjugate of grad.)r   rŠ   rü   r   r   r   Ú	_ConjGradŸ  rp   r¢  Ú
ComplexAbsc              
   C   s>   t  t  |t |¡¡| jd  t  | jd t | jd ¡¡¡S )z#Returns the gradient of ComplexAbs.r   )r   r#  r•  r	   rÂ   r%   r!   r   r   r   r   Ú_ComplexAbsGrad¥  s   
ÿÿÿýr¤  ÚCastc                 C   sR   t jt jt jt jt jt jg}| jd jj	}|jj	}||v r'||v r't
 ||¡S d S r÷   )r   Úfloat16r	  Úfloat64Úbfloat16Ú	complex64Ú
complex128r%   rH   r
  r   re   )r   r   rK   Úsrc_typeÚdst_typer   r   r   Ú	_CastGrad¯  s   þr­  ÚCrossc                 C   s,   | j d }| j d }t ||¡t ||¡fS r-  )r%   r   Úcross)r   r   ÚuÚvr   r   r   Ú
_CrossGrad½  s   

r²  ÚCumsumc                 C   s6   | j d }|  d¡}|  d¡}tj|||| dd gS )Nr   r‚   rƒ   ©r‚   rƒ   )r%   r"   r   Úcumsum)r   r   r   r‚   rƒ   r   r   r   Ú_CumsumGradÄ  s   


þr¶  ÚCumprodc                 C   sb   | j d }| j d }|  d¡}|  d¡}tj||||d}tj|| ||| d}t ||¡d gS )Nr   r   r‚   rƒ   r´  )r%   r"   r   r‰   rµ  r#  )r   r   r   r   r‚   rƒ   rt   r™   r   r   r   Ú_CumprodGradÏ  s   



ÿr¸  ÚCumulativeLogsumexpc                 C   sÄ   | j d }| j d }| jd }|  d¡}|  d¡}t t |d¡t |¡|jj	¡}t t 
|d¡t | ¡|jj	¡}t tj|| || |d| ¡}	t tj|| || |d| ¡}
|	|
 d gS )Nr   r   r‚   rƒ   )r   rƒ   r‚   )r%   r!   r"   r	   rÃ   r   rÐ   r  rH   ÚminÚlessr1  Úcumulative_logsumexp)r   r   r   r   r¼  r‚   rƒ   Úlog_grad_positiveÚlog_grad_negativeÚ
output_posÚ
output_negr   r   r   Ú_CumulativeLogsumexpGradÜ  s@   





ý

ýþþÿþþÿrÁ  Ú	NextAfterc           
      C   s¸   | j d }| j d }t |¡}t |¡}t ||¡\}}t |g¡0 tj||jd}tj	||jd}	t 
t || |¡|¡t 
t |	| |¡|¡fW  d  ƒ S 1 sUw   Y  dS )z@Returns gradient of nextafter(x1, x2) with respect to x1 and x2.r   r   rG   N)r%   r	   r$   r
   r/   r   r  rž   rH   rÄ   r&   r   rg   )
r   r   Úx1r`  Ús_x1Ús_x2Úr_x1Úr_x2Ú
partial_x1Ú
partial_x2r   r   r   Ú_NextAfterGradþ  s    



ÿÿþ$ýrÊ  r  )Ú__doc__ÚnumpyrP   Útensorflow.python.compatr   Útensorflow.python.eagerr   Útensorflow.python.frameworkr   r   r   r   r   Útensorflow.python.opsr	   r
   r   r   r   r   ÚRegisterGradientr   r   r*   rC   rD   rE   rd   rk   ro   rr   r   rš   rœ   r£   r¯   r±   r¶   rº   r¿   rÁ   rÉ   rÌ   rÎ   rÞ   ré   rë   rí   rð   rò   rõ   rú   rý   r  r  r  r  r  r  r#  r&  r+  r.  r2  r7  r:  rD  rH  rK  rN  rQ  rU  rY  r[  rc  rf  rm  rq  ru  rx  r{  r  r  rƒ  r‡  r‹  r  r’  r–  r  r¡  r¥  r©  r¬  r¯  r²  rµ  r¸  r»  rÂ  rÅ  rÈ  rÐ  rÕ  rØ  rÛ  rÝ  rß  rá  rå  rê  rì  rï  rö  rø  rü  r  r  r  r  r  r  r  r  r!  r$  r,  r2  r3  r6  r:  r>  ÚNotDifferentiablerK  rP  rY  r[  r]  ro  rr  ru  rw  ry  rƒ  rŒ  r’  r˜  r›  r   r¢  r¤  r­  r²  r¶  r¸  rÁ  rÊ  r   r   r   r   Ú<module>   sB  


G
G



0










)ý#



/




	
	
	





	












	





	







	



	

























!
!
(






4

#

(



 
3





,






	





!