
    ij>                        d dl Z d dlmc 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 Zd Zd Zd Zd Zd Zd Zd ZdNdZd Zd ZdOdZd Z dPdZ!d Z"d Z#dQdZ$d Z%dRdZ&dQdZ'dSd Z(d! Z)dNd"Z*d# Z+dSd$Z,dSd%Z-dSd&Z.	 dTd'Z/	 dTd(Z0d) Z1d* Z2d+ Z3d, Z4d- Z5	 	 	 dUd.Z6	 	 	 dUd/Z7dVd0Z8dVd1Z9	 	 	 	 dWd2Z:	 	 	 	 dWd3Z;	 	 	 	 dWd4Z<	 	 	 	 	 dXd5Z=dYd6Z>dYd7Z?dZd8Z@dZd9ZAd[d:ZBd\d;ZC	 d]d<ZDd^d=ZE	 	 d_d>ZFd?ZGd@ ZHdA ZI	 	 	 d`dBZJdC ZKdD ZL	 	 	 	 	 dadEZMdF ZNdG ZO	 dbdHZP	 	 	 	 	 	 dcdIZQdddJZRdddKZSdedLZTdedMZUy)f    N)backend)check_conv_input_channels)#check_conv_transpose_input_channels)-compute_conv_transpose_output_crops_for_torch)cast)convert_to_tensor)
get_device)expand_dims)where)standardize_tuplec                 B    t        |       } t        j                  |       S N)r   tnnreluxs    o/var/www/html/emotional.easysim.app/public_html/venv/lib/python3.12/site-packages/keras/src/backend/torch/nn.pyr   r          !A88A;    c                 B    t        |       } t        j                  |       S r   )r   r   relu6r   s    r   r   r      s    !A99Q<r   c                 B    t        |       } t        j                  |       S r   )r   r   sigmoidr   s    r   r   r      s    !A;;q>r   c                 0   t        |       } t        j                  | dk  t        j                  d| j                  | j
                        t        j                  | dk\  t        j                  d| j                  | j
                        d| dz   z              S )N        devicedtype         ?      ?)r   torchr   tensorr   r   r   s    r   sparse_sigmoidr%   #   sr    !A;;	RS9FLLQXXQWW=1q5M	
 r   c                 B    t        |       } t        j                  |       S r   )r   r   tanhr   s    r   r'   r'   0   r   r   c                 B    t        |       } t        j                  |       S r   )r   r   
tanhshrinkr   s    r   tanh_shrinkr*   5       !A>>!r   c                 B    t        |       } t        j                  |       S r   )r   r   softplusr   s    r   r-   r-   :       !A<<?r   c                 B    t        |       } t        j                  |       S r   )r   r   softsignr   s    r   r0   r0   ?   r.   r   c                 F    t        |       } t        j                  | |      S N)lambd)r   r   
softshrinkr   	thresholds     r   soft_shrinkr7   D       !A>>!9--r   c           
          t        |       } t        j                  | dk  t        j                  |       t        j                  | dk  d| dz   dz  z  |             S )Nr   r    g      ?   )r   r#   r   
zeros_liker   s    r   sparse_plusr<   I   sS    !A;;	RAEEa!e\115 r   c                 B    t        |       } t        j                  |       S r   )r   r   silur   s    r   r>   r>   R   r   r   c                 t    t        |       } t        |      }| t        j                  | dz  |z         z   }|dz  S )Nr:   )r   r#   sqrt)r   bys      r   
squareplusrC   W   s:    !A!A	EJJq!tax  Aq5Lr   c                 B    t        |       } t        j                  |       S r   )r   r   
logsigmoidr   s    r   log_sigmoidrF   ^   r+   r   c                 F    t        |       } t        j                  | |      S )N)negative_slope)r   r   
leaky_relu)r   rH   s     r   rI   rI   c   s    !A>>!N;;r   c                 B    t        |       } t        j                  |       S r   )r   r   hardsigmoidr   s    r   hard_sigmoidrL   h   s    !A??1r   c                 B    t        |       } t        j                  |       S r   )r   r   	hardswishr   s    r   	hard_silurO   m   s    !A==r   c                 D    t        |       } t        j                  | |      S r   )r   r   elur   alphas     r   rQ   rQ   r   s    !A771er   c                 B    t        |       } t        j                  |       S r   )r   r   selur   s    r   rU   rU   w   r   r   c                 t    t        |       } |rt        j                  | d      S t        j                  |       S )Nr'   )approximate)r   r   gelu)r   rW   s     r   rX   rX   |   s.    !Axxv..88A;r   c                 F    t        |       } t        j                  | |      S )N)rS   )r   r   celurR   s     r   rZ   rZ      s    !A88AU##r   c                 F    t        |       } t        j                  | |      S )Ndim)r   r   glu)r   axiss     r   r^   r^      s    !A771$r   c                 H    t        |       } t        j                  | dd      S )Ng      r!   )min_valmax_val)r   r   hardtanhr   s    r   	hard_tanhrd      s    !A<<455r   c                 F    t        |       } t        j                  | |      S r2   )r   r   
hardshrinkr5   s     r   hard_shrinkrg      r8   r   c                 H    t        |       } t        j                  | ||      S )N)r6   value)r   r   r6   )r   r6   default_values      r   r6   r6      s    !A==i}EEr   c                    t        |       } t        j                  | j                        }t	               dk(  r.t        j                  | j                        dk(  rt        | d      } |Ot        j                  | dg      }t        j                  |d      }t        j                  || j                        }nt        j                  | |      }t        ||      S Ncpufloat16float32r   r\   )r   r   standardize_dtyper   r	   r   r#   reshaper   softmaxshaper   r_   r   outputs       r   rr   rr      s    !A%%agg.E 	%%agg.);I| q2$'V,vqww/QD)r   c                    t        |       } t        j                  | j                        }t	               dk(  r.t        j                  | j                        dk(  rt        | d      } |Ot        j                  | dg      }t        j                  |d      }t        j                  || j                        }nt        j                  | |      }t        ||      S rl   )r   r   rp   r   r	   r   r#   rq   r   log_softmaxrs   rt   s       r   rw   rw      s    !A%%agg.E 	%%agg.);I| q2$'R0vqww/-r   c                 r   t        |       }t        j                  ||d      \  }}t        j                  ||      }t        j                  d|j                  |      dz   |j                  |j                        }dg|j                  z  }d||<   |j                  |      }||dz
  |z  z
  dkD  }t        j                  ||d      }	t        j                  ||t        j                  d	|j                  
            }
t        j                  |
|d      dz
  |	z  }t        j                  ||z
  d	      }|S )NT)r]   
descendingr\   r    r   r   r   r]   keepdimr   r   min)r   r#   sortcumsumarangesizer   r   ndimviewsumr   r$   clamp)r   r_   logitslogits_sorted_logits_cumsumrr_shapesupportklogits_cumsum_safetauru   s                r   	sparsemaxr      s   q!Fzz&dtDM1LLD9M	6;;tq fll	A cFKKGGDM	wA}q0A559G		'tT2AS G 99'T4@1D
IC[[#3/FMr   c                     |dz
  |z  dz   }|dk(  r|dz
  }n#| |z   dz
  |z  }t        d|dz
  |z  |z   | z
        }|dz  }||z
  }||fS )zSCompute padding length along one dimension with support
    for asymmetric padding.r    r   r:   )max)	input_lengthkernel_lengthstridedilation_rateeffective_k_sizetotal_paddingoutput_sizeleft_paddingright_paddings	            r   _compute_padding_lengthr      s    
 &)]:Q>{(1, $f,q0V;a6),<<|K

 !A%L!L0M-((r   c                    | j                   dd }t        |      }g }|dk7  rt        ||d      }t        |      D ]6  }	|dk(  rdn||	   }
t	        ||	   ||	   ||	   |
      }|j                  |       8 t        d |D              r| |D cg c]  \  }}|	 c}}fS g }t        |      D ]  }|j                  |        |dk(  rdnd}t        j                  | t        |      |	      d
fS c c}}w )aZ  Apply same padding to the input tensor.

    This function will evaluate if the padding value is compatible with torch
    functions. To avoid calling `pad()` as much as possible, which may cause
    performance or memory issues, when compatible, it does not apply the padding
    to the tensor, but returns the input tensor and the padding value to pass to
    the torch functions. If not compatible, it returns the padded tensor and 0
    as the padding value.

    Returns:
        tensor: A padded tensor or the inputs.
        padding: The padding value, ready to pass to the torch functions.
    r:   Npoolingr   r    c              3   ,   K   | ]  \  }}||k(    y wr    ).0leftrights      r   	<genexpr>z&_apply_same_padding.<locals>.<genexpr>  s     
4[T545=
4s   	replicateconstant)padmoder   )rs   lenr   ranger   appendallreversedextendr   r   tuple)inputskernel_sizestridesdata_formatoperation_typer   spatial_shapenum_spatial_dimspaddingidilr   r   r   flattened_paddingr   s                   r   _apply_same_paddingr      s     LL$M=)G")+_
 #$ !Y.aM!4D%!k!ngaj#
 	s 
4G
44G4q444   &  %& )I5;:D776u%67dCQFF 5s   C+c                 H   | j                   dz
  }|dk(  r$t        j                  | d      j                         S |dk(  r$t        j                  | d      j                         S |dk(  r$t        j                  | d      j                         S t	        d| j
                   d      )	z=Transpose inputs from channels_last to channels_first format.r:   r    r   r:   r    )r      r    r:   r   )r      r    r:   r   z^Inputs must have ndim=3, 4 or 5, corresponding to 1D, 2D and 3D inputs. Received input shape: .)r   r#   permute
contiguous
ValueErrorrs   )r   r   s     r   _transpose_spatial_inputsr     s     ;;?Dqy}}VY/::<<	}}V\2==??	}}V_5@@BB
	!!'a	1 r   c                     t        | j                        dz
  }|dk(  rt        j                  | d      } | S |dk(  rt        j                  | d      } | S |dk(  rt        j                  | d      } | S )Nr:   r    r   )r   r:   r   r    r   )r   r:   r   r   r    r   rs   r#   r   )outputsr   s     r   _transpose_spatial_outputsr   4  su    7==)A-1--3
 N	 
Q	--6 N 
Q	--9Nr   c                     t        | j                        dz
  }|dk(  rt        j                  | d      } | S |dk(  rt        j                  | d      } | S |dk(  rt        j                  | d      } | S )Nr:   r    )r:   r    r   )r   r:   r   r    r   )r   r   r   r    r:   r   )kernelr   s     r   _transpose_conv_kernelr   @  sw     6<<(1,1vy1
 M	 
Q	v|4 M 
Q	v7Mr   c                 X    | dk(  rt         j                  S | dk(  rt         j                  S y )Nr      )r#   channels_lastchannels_last_3d)r   s    r    _get_channels_last_memory_formatr   M  s+    qy"""	%%%r   c                 |    t        | j                        }|$| j                  |      s| j                  |      S | S )N)memory_format)r   r   is_contiguousr   )r$   mem_fmts     r   _maybe_convert_to_channels_lastr   U  s?    .v{{;G6#7#7g#7#N  w 77Mr   c                    t        |       } | j                  dz
  }t        ||d      }||}nt        ||d      }t        j                  |      }|dk(  rt        |       } |dk(  rt        | |||d      \  } }nd}t               }|dk(  r,t        j                  | j                  | j                  d	
      } |dk(  rt        j                  | |||      }nW|dk(  rt        j                  | |||      }n8|dk(  rt        j                  | |||      }nt!        d| j                   d      |j#                  |      }|dk(  rt%        |      }|S )z!Fixed max pooling implementation.r:   	pool_sizer   r   samer   r   metarm   r   r   r   r    )r   r   r   r   lInputs to pooling op must have ndim=3, 4 or 5, corresponding to 1D, 2D and 3D inputs. Received input shape: r   )r   r   r   r   standardize_data_formatr   r   r	   r#   emptyrs   r   r   
max_pool1d
max_pool2d
max_pool3dr   tor   )r   r   r   r   r   r   r   r   s           r   max_poolr   \  sl    v&F{{Q!)-={KI#G-=yI11+>Ko%*62& .IwY
 \F V\\%
 1..	'7
 
Q	..	'7
 
Q	..	'7
 %%+\\N!5
 	
 jj Go%,W5Nr   c                 
   t        |       } | j                  dz
  }t        ||d      }||nt        ||d      }t        j                  |      }|}|dk(  rt        |       } |dk(  rt        | ||dd      \  } }nd}|d	k(  rt        j                  | |||d
      }nY|dk(  rt        j                  | |||d
      }n9|dk(  rt        j                  | |||d
      }nt        d| j                   d      |dk(  rt        |      }|S )z7Fixed average pooling with correct padding calculation.r:   r   r   r   r   channels_firstr   r   r    F)r   r   r   count_include_padr   r   r   )r   r   r   r   r   r   r   r   
avg_pool1d
avg_pool2d
avg_pool3dr   rs   r   )r   r   r   r   r   r   orig_formatr   s           r   average_poolr     sR    v&F{{Q!)-={KI ? 	w(8)D  11+>KKo%*62& .
  1..!#
 
Q	..!#
 
Q	..!#
 %%+\\N!5
 	
 o%,W5Nr   c                 P   t        |       } | j                  dz
  }t        j                  |      }|}|dk(  rt	        |       } t        |t              r|dk(  r|n|f|z  }nt        ||d      }t               dk(  r,t        j                  | j                  | j                  d      } |dk(  rt        j                  | |      }nS|dk(  rt        j                  | |      }n6|d	k(  rt        j                   | |      }nt#        d
| j                   d      |dk(  rt%        |      }|S )z>Adaptive average pooling(1D/2D/3D) with channels_last support.r:   r   r    r   r   rm   r   r   r   zSInputs to adaptive average pooling must have ndim=3, 4 or 5, Received input shape: r   )r   r   r   r   r   
isinstanceintr   r	   r#   r   rs   r   r   adaptive_avg_pool1dadaptive_avg_pool2dadaptive_avg_pool3dr   r   )r   r   r   r   r   torch_output_sizer   s          r   adaptive_average_poolr     s6   v&F{{Q11+>KKo%*62+s#  1$ "22 	 .)=
 |vV\\%
 1))&>OP	Q	))&>OP	Q	))&>OP%%+\\N!5
 	

 o%,W5Nr   c                 ~   t        |       } | j                  dz
  }t        j                  |      }|}|dk(  rt	        |       } t        |t              r|dk(  r|n|f|z  }nt        ||d      }t               dk(  r,t        j                  | j                  | j                  d      } |dk(  rt        j                  | |      }nS|dk(  rt        j                  | |      }n6|d	k(  rt        j                   | |      }nt#        d
| j                   d      t        |t$              r|d   n|}|dk(  rt'        |      }|S )z:Adaptive max pooling(1D/2D/3D) with channels_last support.r:   r   r    r   r   rm   r   r   r   zOInputs to adaptive max pooling must have ndim=3, 4 or 5, Received input shape: r   r   )r   r   r   r   r   r   r   r   r	   r#   r   rs   r   r   adaptive_max_pool1dadaptive_max_pool2dadaptive_max_pool3dr   r   r   )r   r   r   r   r   r   resr   s           r   adaptive_max_poolr     sJ   v&F{{Q11+>KKo%*62+s#  1$ "22 	 .)=
 |vV\\%
 1%%f:KL	Q	%%f:KL	Q	%%f:KL%%+\\N!5
 	

 #3.c!fCGo%,W5Nr   c                    t        |       } t        |      }| j                  dz
  }t        ||d      }t        j                  |      }|dk(  rt        |       } t        |      }|dk(  rt        |       } t        |      }| j                  d   }|j                  d   }||z  dk7  rt        d| d| d      ||z  }	|d	k(  r!t        | |j                  dd
 ||d|      \  } }nd}|dk(  rt        j                  | |||||	      }
n[|dk(  rt        j                  | |||||	      }
n:|dk(  rt        j                  | |||||	      }
nt        d| j                   d      |dk(  rt        |
      }
|
S )z&Convolution with fixed group handling.r:   r   r   r    r   zInput channels (z.) must be divisible by kernel input channels ()r   Nconv)r   r   dilationgroupsr   zrInputs to conv operation should have ndim=3, 4, or 5,corresponding to 1D, 2D and 3D inputs. Received input shape: r   )r   r   r   r   r   r   r   r   rs   r   r   r   conv1dconv2dconv3dr   )r   r   r   r   r   r   r   in_channelskernel_in_channelsr   r   s              r   r   r   <  s    v&Fv&F{{Q)99EG11+>Ko%*62#F+Fo%0808 ,,q/Ka''1,{m ,&&8%9<
 	
 ..F &-LL
  1**"
 
Q	**"
 
Q	**"
 ll^1&
 	
 o%,W5Nr   c                    t        j                  |      }t        |       } t        |      }t        | ||       t	        j
                  ||j                  d d d|j                  d   |j                  d   z  fz         }t        | |||||      S )Nr    r   )r   r   r   r   r#   rq   rs   r   )r   r   r   r   r   r   s         r   depthwise_convr    s     11+>Kv&Fv&Fffk:]]Sb!QR(86<<;K(K$LLF +}MMr   c                     t        j                  |      }t        |       } t        |      }t        |      }t        | ||       t	        | |||||      }t        ||dd||      S )Nr    valid)r   r   r   r   )r   r   r   r   r  r   )r   depthwise_kernelpointwise_kernelr   r   r   r   depthwise_conv_outputs           r   separable_convr    s     11+>Kv&F()9:()9:f&6D* # r   c                    t        |       } t        |      }| j                  dz
  }t        ||d      }t        j                  |      }t        | ||       t        | j                  |j                  ||||      }|dk(  rt        |       } t        |      }t        |t              r|g|z  }|dk(  rt        j                  | ||dd|      }	n[|dk(  rt        j                  | ||dd|      }	n:|dk(  rt        j                  | ||dd|      }	nt!        d	| j                   d
      t#        d       t#        d       g}
t%        d |D              }|D ]6  \  }}t'        d|      }|dkD  r| nd }|
j)                  t#        ||             8 |	t+        |
         }	|rNg }t-        |      D ](  \  }}|j/                  |dk  r| nd|dk  r| ndg       * t        j0                  |	|      }	|dk(  rt3        |	      }	|	S )Nr:   r   )input_shapekernel_shaper   r   output_paddingr   r   r    r   )r   r   r  r   r   z|Inputs to conv transpose operation should have ndim=3, 4, or 5,corresponding to 1D, 2D and 3D inputs. Received input shape: r   c              3   :   K   | ]  \  }}|d k  xs |d k    yw)r   Nr   )r   clcrs      r   r   z!conv_transpose.<locals>.<genexpr>  s$     >fb"a)26)>s   )r   r   r   r   r   r   r   rs   r   r   r   r   r   conv_transpose1dconv_transpose2dconv_transpose3dr   sliceanyr   r   r   r   r   r   r   )r   r   r   r   r  r   r   r   cropsr   slicesneeds_zero_pad	crop_left
crop_rightstartendpadss                    r   conv_transposer"    s<    v&Fv&F{{Q)99EG11+>K'D :LL\\%#E o%*62#F+F-%&*::1&&"
 
Q	&&"
 
Q	&&"
 ll^1&
 	
 Dk5;'F>>>N!& )	:Ay!'!^zkeE3'() eFm$G%-e_ 	!IzKK"+a-YJQ#->ZKq	 '''4(o%,W5Nr   c                    |rt        d      t        | t        j                        } t        dt        j                        }t	        j
                  t        j                  | d      |      }t        t        | d      dk\  ||      }t        ||      }|j                         }|dk7  rP||k7  rKt        t        |            }d||<   t        |dz   |      D ]  }	||	xx   dz  cc<    |j                  |      }|S )N2Unsupported value `sparse=True` with torch backendr   r   r}   r   r_   r    )r   r   r#   longr   one_hotr   r   r
   r]   listr   r   )
r   num_classesr_   r   sparsezeroru   dimsnew_axes_orderaxs
             r   r(  r(  +  s    MNN 	!5::.AQejj1D [[QA.<F;qr*a/>FvU3F::<DrzddleDk*!tq$' 	$B2!#	$/Mr   c                     |rt        d      t        |       } t        | j                        dkD  rdnd}t	        j
                  t        t        | d      |||      |      }|S )Nr$  r    r   int32)r_   r   r\   )r   r   r   rs   r#   amaxr(  r   )r   r*  r_   r   r+  reduction_axisr   s          r   	multi_hotr4  E  s^    MNN!Aagg,*QNjjQ +DFG Nr   c                 j   t        |       } t        |      }| j                  |j                  k7  r%t        d| j                   d|j                         t        | j                        dk  r%t        d| j                   d|j                         |rt	        j
                  ||      }nn|t        j                  ||d      z  }t        j                  |t        j                         dt        j                         z
        }t        j                  |      }t        j                  | |z  |       S )	NQArguments `target` and `output` must have the same shape. Received: target.shape=, output.shape=r    zPArguments `target` and `output` must be at least rank 1. Received: target.shape=r\   Trz   r!   )r   rs   r   r   r   rw   r#   r   clipr   epsilonlog)targetru   from_logitsr_   log_probs        r   categorical_crossentropyr>  Q  s   v&Fv&F||v||#"LL>H
 	

 6<<1"LL>H
 	
 ??6t4%))FdCCFGOO$5sW__=N7NO99V$IIfx'T222r   c                 4   t        | t        j                        } t        |      }t        | j                        t        |j                        k(  r)| j                  |   dk(  rt        j
                  | |      } t        |j                        dk  rt        d|j                         t        |j                        }||= t        | j                        |k7  r%t        d| j                   d|j                         |j                         dk(  r%|j                  d      }| j                  d      } d}n,d	}||j                         z  }|dk7  r|j                  |d      }|rt        j                  || d
      }n|t        j                  |dd      z  }t        j                  |t        j                          dt        j                          z
        }t        j"                  |      }t        j$                  || d
      }|r|j                  d      }|S )Nr%  r    r\   zBArgument `output` must be at least rank 1. Received: output.shape=zcArguments `target` and `output` must have the same shape up until the last dimension: target.shape=r7  r   TFnone	reductionrz   r!   )r   r#   r'  r   rs   squeezer   r)  r]   	unsqueezemovedimr   cross_entropyr   r8  r   r9  r:  nll_loss)	r;  ru   r<  r_   output_shape_without_class_dimrC  
class_axisresultr=  s	            r   sparse_categorical_crossentropyrK  k  s   vUZZ8Fv&F
6<<C--&,,t2D2Iv40
6<<1"LL>+
 	

 &*&,,%7"&t,FLL;;"LL>H
 	
 zz|q!!!$!!!$FJJL(
?^^J2F""66VD%))F4@@FGOO$5sW__=N7NO99V$h&A"Mr   c                    t        |       } t        |      }t        j                  j                  j	                         r| j
                  dkD  r|j
                  | j
                  k(  rl| j                  d   dk(  rZ|j                  d   dk(  rHt        j                  | d      j                         } t        j                  |d      j                         }| j                  |j                  k7  r%t        d| j                   d|j                         |rt        j                  || d      S t        j                  |t        j                         dt        j                         z
        }t        j                  || d      S )Nr    r   r6  r7  r@  rA  r!   )r   r#   backendsmpsis_availabler   rs   rC  r   r   r    binary_cross_entropy_with_logitsr8  r   r9  binary_cross_entropy)r;  ru   r<  s      r   binary_crossentropyrR    s3   v&Fv&F 	'')KK!OKK6;;&LL!LL!vr*557vr*557||v||#"LL>H
 	
 33Ff
 	
 FGOO$5sW__=N7NO''&IIr   c                    |rt        d      t        |       } d}t        j                  | j                        }|dk(  rd}t        | d      } t        j                  | |d      }t        j                  t        j                  |       |d      t        j                  |      z
  }|s,t        j                  ||      }t        j                  ||      }|rt        j                  |t        j                  t        j                        j                  t        j                  t        j                        j                        }t        j                  |t        j                  t        j                        j                  t        j                  t        j                        j                        }t        ||      }t        ||      }||fS )Nz9Argument synchronized=True is not supported with PyTorch.Frn   Tro   rz   )NotImplementedErrorr   r   rp   r   r   r#   meansquarerC  r8  finforn   r~   r   )r   axeskeepdimssynchronized	need_cast	ori_dtyperU  variances           r   momentsr^    s_   !G
 	
 	!A I))!''2II	I::aT40D zzQT4TH }}T4(==40zzKK&**KK&**

 ::KK&**KK&**

 D)$),>r   c                    t        |       } t        |      }t        |      }dgt        | j                        z  }|j                  d   ||<   t        j                  ||      }t        j                  ||      }|"t        |      }t        j                  ||      }nt        j
                  |      }|"t        |      }t        j                  ||      }nt        j                  |      }| j                  |      j                  |j                  |      j                         j                  |            j                  |      S )Nr    r   )r   r   rs   r#   rq   r;   	ones_likesubtractmul_addrsqrt_muladd_)r   rU  r]  r_   offsetscaler9  rs   s           r   batch_normalizationri    s    	!AT"D *HC#agg,E**Q-E$K==u%D}}Xu-H"6*vu-!!$'!%(eU+) 	


4	hll7#**,007	8	fr   c                 H   t        |       } t        |      }t        |      }t        |      }t        j                  |j                  d      }t	        ||      }t        j                  |dd      }t        j                  |d      }t        j                  || |||d      }|S )Nro   r    r   r   r\   r@  )blankrB  )
r   r   result_typer   r   r#   	transposer   rw   ctc_loss)r;  ru   target_lengthoutput_length
mask_indexr   r   losss           r   rn  rn  
  s    v&Fv&F%m4M%m4M i8E&% F__VQ*F__V,F<<D Kr   c                    t        |       } t        |d      }| j                  \  }}}||dz
  }t        j                  | d      }t	        |d      }t        j
                  | d      d   }t        j                  ||j                        d d d f   }	|	|d d d f   k\  }	t        j                  |	||      }t        j                  |	d|      }|rD|d d dd f   |d d d df   k(  }
t        j                  |
d	      }
t        j                  |
||      }||k(  }t        j                  |d|      }t        j                  t        j                  ||j                        d
      }t        j                  ||df      }t        j                  |||      }t        j                  |d
      }t        j                  ||d
      }t        j                  |d      d d d f    }t        j                  |d
      }||fS )Nr1  r%  r    r   r&  r   r|   r   )r    r   r   r   r\   )r   rs   r#   argmaxr   r   r   r   r   r   r   rD  tileargsorttake_along_dimr   )r   sequence_lengthsmerge_repeatedrq  
batch_size
max_lengthr*  indicesscoresseqlen_maskrepeatinvalid_maskorders                r   _ctc_greedy_decoder  !  s    v&F()9I*0,,'J
K 1_
ll6+G7G$GYYvB'*F,,z'..A$'JK!1!T'!::Kkk+z7;G[[c62FAB71crc6?2.++fj': j(Lkk,G4G OOZ7QE JJuz1o.EKKj%8EMM%R(E""7Er:GiiQ'400Foog1-GF?r   l   yn< c                    | j                   \  }}| j                  }| j                  t        j                        dz   }t
        t        j                  |t        j                  |      z  }||z  j                  d      }t        j                  |d      }	| |	   }
||	   }|dd |dd k(  }|
dd |
dd k(  j                  d      }||z  }t        j                  t        j                  dt        j                  |      | g      }t        j                  |j                  t        j                        d      dz
  }t        j                  |t        j                  |      }|||	<   |
|   }|j                   d   }||k  rAt        j                  ||z
  |f||j                   |      }t        j                  ||gd      }||fS )	a  Hash-based row-dedup, padded to a fixed leading size with `pad` rows.

    Mirrors `jax.numpy.unique(..., size=size, fill_value=pad, axis=0,
    return_inverse=True)` in observable behavior: the unique rows come
    first, followed by `(size - n_unique)` rows filled with `pad`, and
    `inverse` maps each input row to its index in the unique output.

    Internally avoids `torch.unique(dim=0)` (which sorts the full row
    tensor every call and dominates beam-search runtime) by hashing each
    path to int64 and deduplicating along that 1D axis. Collisions are
    detected by an explicit row-equality check on adjacent same-hash
    entries, so correctness does not rely on a collision-free hash.
    r    r   r   r\   TstableNr   r   )rs   r   r   r#   int64_KNUTH_HASH_CONSTANTr   r   rv  r   catonesboolr   r   fullr   )pathsr   r   ntr   ppowershashesr  sorted_pathssorted_hashesadj_hash_eq
adj_row_eqis_dupis_first	cum_firstinverseuniquen_uniquepad_rowss                        r   _unique_paddedr  U  s    ;;DAq\\F
 	!A!U\\	V& F &j!$FMM&.E<L5MM#}Sb'99Kqr"l3B&77<<<CJ:%Fyy	AUZZ	7&AH X[[51=AIkk!5;;v>GGEN(#F||AH$::H_a ,,	
 FH-157?r   c                     |j                         }t        j                  ||z
        }t        j                  ||j                  |j
                        }|j                  d| |       t        j                  |      |z   S )z=Log-space scatter-add of `scores` into `num_uniques` buckets.r  r   )r   r#   expzerosr   r   scatter_add_r:  )r  r}  num_uniques
scores_max
scores_expouts         r   _merge_scoresr    s_    J6J./J
++kfmm
LCQ,99S>J&&r   c                    t        |       } t        |d      }| j                  \  }}}t        j                  | d      } ||dz
  }t	        j
                  | dg      } ||z
  dz
  }d}| j                  }	|j                         j                         j                         }
t        ||      }g }g }t        |      D ]  }| |   }t        |
|         }t	        j                  d|z  |f|t        j                  |		      }t	        j                  |d
   d      | d }t	        j                   ||k(  ||j#                  t        j                              }||d|d
f<   t	        j                  d|z  ft%        d      | j&                  |		      }|d
|f   |d| |}|}|ddd
f   |k(  }t        d|      D ].  }t)        |||||   |||      \  }}}t+        ||||||      \  }}}0 t-        |d|z  |z  |      \  }}t/        |||j                  d
         }t	        j                  |d      | d j                  d
      }|j1                  ||          |j1                  ||           t	        j2                  |d
      }t	        j2                  |d
      }t	        j                   ||k(  |||z
  dz
        }|j5                  dd
d      }||fS )au  Beam search CTC decoding for the torch backend.

    Direct port of `keras/src/backend/jax/nn.py::_ctc_beam_search_decode`.
    The semantics (tie-breaking, log-space score merging, blank/emit
    score tracks) match the JAX reference implementation; correctness is
    prioritized over throughput, so the batch dimension is iterated in
    Python rather than vmapped.
    r1  r%  r   r\   Nr    r:   )r-  r  r   Tr  -infr   r   )r   rs   r   rw   r#   flipr   detachrm   tolistr~   r   r   r  r1  rv  r   r   floatr   _ctc_beam_extend_ctc_beam_pruner  r  r   stackr   )r   rx  
beam_width	top_pathsrq  rz  max_seq_lenr*  _padr   
seqlen_cpunum_init_pathspaths_per_batchscores_per_batchrA   r   	seq_len_b
init_pathsmax_classesinit_classesinit_scoresr  r}  maskedr  paths_uniquer  top_indicess                               r   _ctc_beam_search_decoder    s    v&F()9I+1<<(J[__V,F 1_

 ZZaS)Fz)A-JD]]F "((*..0779Jj1NO: .51I
1&	ZZ^[)++	

 mmAaD67GH{{:%t[^^EKK-H
 *6
?N?A%&jj^&M,,	
 ()K'8O^$q!t$ q)$ 	A$4vvqt[*d%!E66 %4vv{J%!E66		 !/K*4$!
g w0B0B10EFmmF48)EJJ1M|K89{ 34].5` KKQ/E[[)q1F KKt[5-@1-DEEMM!Q"E&=r   c                    | j                  |d      } |j                  |      }|j                  |      }| |k(  }|j                  t        j                        j	                  d      }t        j
                  | j                  d   | j                        }	| |	|dz
  f   }
t        j                  |dk(  ||
      }
t        j
                  || j                  | j                        }|j                         }|||<   |j                  | j                  d   |z        }|}||k(  }| |
|k(  z  }t        j                  |||      }|| |	|f<   ||j                  | j                  d   |z        z   }| ||fS )zAExtend each beam with every possible class for a single timestep.r   r\   r    r|   r   )repeat_interleaver   r#   r1  rt  r   rs   r   r   r   cloner  )r  r}  r  r   r*  rq  r  is_padpath_tail_indexr   
path_tailsclassesprev_maskedmasked_repeats                 r   r  r    s^   ##KQ#7E%%k2F%%k2Fd]Fii,333:O\\%++a.>Fv223J_14DJll;u||5;;OGmmoGGJnnU[[^{:;GK_F!\jG&;<Mkk-w7G%,E&/
!"ahhu{{1~<==F&&  r   c                    t        | d|z  |z  |      \  }}t        j                  |t        d      |      }t        j                  ||t        d            }	|j                  d   }
t        |||
      }t        ||	|
      }	t        j                  ||	      }t        j                  |d      | d }||   }||   }|	|   }|j                  dd      } t        j                  ||g      }t        j                  t        j                  |t        j                  | j                  	      t        j                  |t        j                  | j                  	      g      }| ||fS )
zADedup + score-merge + keep top `beam_width` (emit, blank) tracks.r:   r  r  r   Tr  Nr    r  )r  r#   r   r  rs   r  	logaddexprv  r  r  r  r  r   r  )r  r}  r  r*  r  r  r  r  emit_scoresmask_scores	n_uniquestotal_scoresr  	paths_topemit_scores_topmask_scores_top
masked_outs                    r   r  r    s9   *AOj0dL' ++feFmV<K++ffeFm<K""1%Ii@Ki@K??;<L--T:J;<HK[)I!+.O!+.OQ"EYY9:FKK
%**U\\JJJzELLI	
J &*$$r   c                     t        |       } t        j                  | j                  d      }t	        | |      } |dk(  rt        | |||      S |dk(  rt        | ||||      S t        d| d      )Nro   greedy)ry  rq  beam_search)r  r  rq  zInvalid strategy z2. Supported values are 'greedy' and 'beam_search'.)r   r   rl  r   r   r  r  r   )r   rx  strategyr  r  ry  rq  r   s           r   
ctc_decoder  0  s     v&Fi8E&% F8!)!	
 	
 
]	"&!!
 	
 z ** *
 	
r   c                 v   | j                   |j                   k7  r&t        d| j                    d|j                    d      t        |       t        |      }} t        || j                        }t	        j
                  | |z
  dz        }dt	        j                  |      z  dt	        j                  |      z  z
  }|S )NzInput shapes z and z" must match for PSNR calculation. r%  r:      
   )rs   r   r   r   r#   rU  log10)x1x2rb   msepsnrs        r   r  r  S  s    	xx288BHH:U288* 5+ +
 	
 	"" 	B  rxx8G
**b2g!^
$CG$$rEKK,<'<<DKr   c                 `    t        j                  |       } | dk(  rd}nd}t        |dz  |       S )Nrn   g    @g̓$Ggffffffr%  )r   rp   r   )r   vals     r   _get_large_negativer  d  s5    %%e,E	S4Zu55r   c           	          	 ddl m} ddl m} 	  || |||d|d      }|r ||d      du rt        d       ||d      S # t        $ r |rt        d      Y yw xY w# t        $ r  || |||d|      }Y Uw xY w)	z+Verify the availability of flash attention.r   )
SDPAParams)can_use_flash_attentionzFlash attention is not supported in your current PyTorch version. Please update it by following the official guide: https://pytorch.org/get-started/locally/Fr   TzfFlash attention is not supported with the provided inputs. Please check the warnings for more details.)torch.backends.cudar  r  ImportError	TypeErrorRuntimeError)	querykeyri   mask	is_causalraise_errorr  r  spda_paramss	            r   _can_use_flash_attentionr  m  s    
2?
 
& .{DAUJ:
 	
 #;66E  ; 
 &  	
 
	
s    > A AAA32A3c	           	      |   t        |       } t        |      }t        |      }t        | j                        dk7  s0t        |j                        dk7  st        |j                        dk7  r3t        d| j                   d|j                   d|j                   d      ||t        d      t	        j
                  | j                  |j                  |j                        }	t        | |	      } t        ||	      }t        ||	      }||nt        |d      }||ry| j                  d	   |j                  d	   }}
t        j                  t        j                  |
|ft        j                  |j                  
            }t        j                  ||      }d}t        j                  |dt        | j                              }|t        ||	      }|}d\  }}t        j                   | ||      } t        j                   |||      }t        j                   |||      }| j                  d	   }|j                  d	   }||kD  r:|d	kD  r5||z  }t        j"                  ||d	      }t        j"                  ||d	      }|t%        | ||||      }n|du rt%        | ||||d       |rt        j&                  j(                  j+                  t        j&                  j(                  j,                  j.                  g      5  t        j&                  j0                  j3                  | |||||      }d d d        nk||j5                         }t        j&                  j0                  j3                  | j5                         |j5                         |j5                         |||      }t        j                   ||      S # 1 sw Y    xY w)Nr   zG`dot_product_attention` only supports 4D inputs. Received: query.shape=z, key.shape=z, value.shape=r   z=Only one of `bias` and `mask` can be provided. Received both.r  r%  r    r  Fr   )r    r:   )repeatsr]   T)r  )rM  )	attn_maskr  rh  )r   r   rs   r   r   rl  r   r   r#   trilr  r  r   logical_andr   r  rm  r  r  nn	attentionsdpa_kernel
SDPBackendFLASH_ATTENTION
functionalscaled_dot_product_attentionr   )r  r  ri   biasr  rh  r  flash_attentionattn_logits_soft_capcompute_dtypeq_lenkv_lencausal_maskaxis0axis1num_query_headsnum_kv_headsr   attention_outputs                      r   dot_product_attentionr    sy    e$E
C
 Ce$E
5;;1CII! 3s5;;7G17L%%*[[Mcii[ I ;;-q*
 	

 D,K
 	
 ''SYYLM&E
sM
"C&E<4%6t6%JD
 "KKNCIIaL6E**

FO5::dkkK
 $$T;7D	{{4&9%++&FG ];LE5OOE5%0E
//#ue
,COOE5%0Ekk!nO99Q<L%,*: L0%%c6qA''v1E23tY
 
D	  	!3tYD	
 XX++hh((33CCD , 
 
	  %xx22OO#  P  
	 
	 ??$D 88..KKNN L 
 ??+UE::-
	 
	s   70N22N;c                 6    t        j                  | ||||      S )a  Native PyTorch implementation of Unfold.
    Extract sliding local blocks from a **NCHW** batched image tensor.

    Args:
        input: 4-D tensor, shape (N, C, H, W)  **required**.
        kernel_size: int or (kH, kW)
        dilation: int or (dH, dW), default 1
        padding: int or (pH, pW), default 0
        stride: int or (sH, sW), default 1

    Returns:
        3-D tensor, shape (N, C*kH*kW, L)
    )r   r   r   r   )r   unfold)inputr   r   r   r   s        r   r  r    s$     :: r   c                 8    t        j                  | |||||      S )a  Native PyTorch implementation of Fold.
    Combine an array of sliding local blocks into a large tensor (col2im).

    Args:
        x: 3-D tensor, shape (N, C*kH*kW, L)  **required**.
        output_size: int or (oH, oW)
        kernel_size: int or (kH, kW)
        dilation: int or (dH, dW), default 1
        padding: int or (pH, pW), default 0
        stride: int or (sH, sW), default 1

    Returns:
        4-D tensor, shape (N, C, oH, oW)
    )r   r   r   r   r   )r   fold)r   r   r   r   r   r   s         r   r  r    s'     88	 r   c                    t        |       } |dk(  ra| j                  \  }}}}||dz  z  }| j                  ||||||      } | j                  dddddd      } | j                  |||z  ||z  |      } | S | j                  \  }}}}||dz  z  }| j                  ||||||      } | j                  dddddd      } | j                  ||||z  ||z        } | S )aQ  PyTorch implementation of depth_to_space.

    Rearranges data from depth into blocks of spatial data.
    Matches TensorFlow's depth_to_space behavior.

    Args:
        x: 4-D tensor with shape (N, H, W, C) for channels_last or
            (N, C, H, W) for channels_first.
        block_size: An integer specifying the block size.
        data_format: "channels_last" or "channels_first".

    Returns:
        A tensor with shape (N, H*block_size, W*block_size, C/block_size**2)
        for channels_last or (N, C/block_size**2, H*block_size, W*block_size)
        for channels_first.
    r   r:   r   r    r   r   r   r   rs   rq   r   )r   
block_sizer   r  hwcnew_cs           r   depth_to_spacer  +  s    " 	!Ao%WW
1aj!m$IIaAz:u=IIaAq!Q'IIaZZ? H WW
1aj!m$IIa
J1=IIaAq!Q'IIaJJ?Hr   c                    t        |       } |dk(  rc| j                  \  }}}}||z  }||z  }| j                  ||||||      } | j                  dddddd      } | j                  |||||dz  z        } | S | j                  \  }}}}||z  }||z  }| j                  ||||||      } | j                  dddddd      } | j                  |||dz  z  ||      } | S )aG  PyTorch implementation of space_to_depth.

    Rearranges blocks of spatial data into depth.
    Matches TensorFlow's space_to_depth behavior.

    Args:
        x: 4-D tensor with shape (N, H, W, C) for channels_last or
            (N, C, H, W) for channels_first.
        block_size: An integer specifying the block size.
        data_format: "channels_last" or "channels_first".

    Returns:
        A tensor with shape (N, H/block_size, W/block_size, C*block_size**2)
        for channels_last or (N, C*block_size**2, H/block_size, W/block_size)
        for channels_first.
    r   r   r    r   r:   r   r   r  )	r   r  r   r  r  r  r  new_hnew_ws	            r   space_to_depthr  T  s   " 	!Ao%WW
1aZZIIa
E:qAIIaAq!Q'IIaq:q='89 H WW
1aZZIIaE:ujAIIaAq!Q'IIaZ]*E59Hr   )r"   )r   )g?)r!   )T)r   )r    )Nr	  Nr   )r    r	  Nr    )r    r	  NNr    )r   NF)Fr   )F)FF)NNgMbP?)r   )TN)d   r    N)r  r  r    Tr   )NFF)NNNFNN)r    r   r    )r   )Vr#   torch.nn.functionalr  r  r   	keras.srcr   &keras.src.backend.common.backend_utilsr   r   r   keras.src.backend.torch.corer   r   r	   keras.src.backend.torch.numpyr
   r   #keras.src.utils.argument_validationr   r   r   r   r%   r'   r*   r-   r0   r7   r<   r>   rC   rF   rI   rL   rO   rQ   rU   rX   rZ   r^   rd   rg   r6   rr   rw   r   r   r   r   r   r   r   r   r   r   r   r   r   r  r  r"  r(  r4  r>  rK  rR  r^  ri  rn  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r   r   r   <module>r%     s=    ! !  L . : 3 5 / A







.


<




$
 
6
.
F
((. 89). NO*GZ,	
 ;B GT(V*` Tt N, F cL4	34.b JF*\ ?C<4 	+b " 3l' ^B!8%D  
F"6 @E)7` 
	
_;D.2&R(r   