
    ij                        d dl Z d dlZd dl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  e j                  d	      d
        Zd Zd Zd Zd Zd Zd Zd Zd ZdPdZd Zd ZdQdZd ZdRdZd Zd Z dSdZ!d Z"dTdZ#dSdZ$dUdZ%d  Z&dPd!Z'd" Z(dUd#Z)dUd$Z*dUd%Z+d& Z,d' Z-	 	 	 dVd(Z.	 	 	 dVd)Z/d* Z0dWd+Z1dWd,Z2dWd-Z3dWd.Z4dWd/Z5dWd0Z6dXd1Z7dXd2Z8d3 Z9	 	 	 	 dYd4Z:d5 Z;d6 Z<	 dXd7Z=	 	 	 	 dYd8Z>	 	 	 	 dYd9Z?	 	 	 	 	 dZd:Z@d[d;ZAd[d<ZBd= ZCd\d>ZDd\d?ZEd]d@ZFd^dAZGdB ZHdC ZI	 d_dDZJd`dEZK	 	 	 	 	 dadFZLdG ZMdH ZNdI ZOdJ ZP	 	 	 	 	 	 dbdKZQdcdLZRdcdMZSdddNZTdddOZUy)e    N)backend)check_conv_input_channels)#check_conv_transpose_input_channels)%compute_adaptive_pooling_window_sizes)#compute_conv_transpose_output_shape)cast)convert_to_tensor)maxsizec                  ^    t        d t        j                  j                         D              S )zWhether the TF runtime sees only CPU devices.

    Cached because the device list does not change during a run and
    `tf.config.list_logical_devices()` is comparatively expensive to call on
    every conv invocation in eager mode.
    c              3   :   K   | ]  }|j                   d k(    yw)CPUN)device_type.0ds     t/var/www/html/emotional.easysim.app/public_html/venv/lib/python3.12/site-packages/keras/src/backend/tensorflow/nn.py	<genexpr>z_cpu_only.<locals>.<genexpr>   s     P!q}}%Ps   )alltfconfiglist_logical_devices     r   	_cpu_onlyr      s"     Pryy/M/M/OPPPr   c                 @    t         j                  j                  |       S N)r   nnreluxs    r   r   r   !       55::a=r   c                 @    t         j                  j                  |       S r   )r   r   relu6r   s    r   r#   r#   %   s    55;;q>r   c                 V    | }t         j                  j                  |       }||_        |S r   )r   r   sigmoid_keras_logits)r    logitsoutputs      r   r%   r%   )   s&    FUU]]1F!FMr   c                    t        |       } t        j                  | dk  t        j                  d| j                        t        j                  | dk\  t        j                  d| j                        d| dz   z              S )N        dtype         ?      ?)r	   r   whereconstantr-   r   s    r   sparse_sigmoidr3   0   sa    !A88	R
Cqww'
aS8#Q-H r   c                 @    t         j                  j                  |       S r   )r   r   tanhr   s    r   r5   r5   9   r!   r   c                 F    | t         j                  j                  |       z
  S r   )r   mathr5   r   s    r   tanh_shrinkr8   =   s    rww||Ar   c                 @    t         j                  j                  |       S r   )r   r7   softplusr   s    r   r:   r:   A   s    77Ar   c                 @    t         j                  j                  |       S r   )r   r   softsignr   s    r   r<   r<   E   s    55>>!r   c                     t        j                  | |kD  | |z
  t        j                  | | k  | |z   t        j                  |                   S r   )r   r1   
zeros_liker    	thresholds     r   soft_shrinkrA   I   sD    88	I	I
iZYa0@A r   c                     t        j                  | dk  t        j                  |       t        j                  | dk  dt        j                  | dz   d      z  |             S )Nr*   r.   g      ?   )r   r1   r>   powr   s    r   sparse_plusrE   Q   sN    88	R
a
Q"&&Q"22A6 r   c                 @    t         j                  j                  |       S r   )r   r   silur   s    r   rG   rG   Y   r!   r   c                     t        |       } t        || j                        }| t        j                  t        j                  |       |z         z   }|dz  S )Nr,   rC   )r	   r-   r   sqrtsquare)r    bys      r   
squareplusrM   ]   sE    !A!177+A	BGGBIIaL1$%%Aq5Lr   c                 @    t         j                  j                  |       S r   )r   r7   log_sigmoidr   s    r   rO   rO   d   s    77q!!r   c                 D    t         j                  j                  | |      S )N)alpha)r   r   
leaky_relu)r    negative_slopes     r   rR   rR   h   s    55A^44r   c                     t        |       } t        | t        j                  d| j                        z         t        j                  d| j                        z  S )Ng      @g      @)r	   r#   r   r2   r-   r   s    r   hard_sigmoidrU   l   s?    !AR[[agg../"++c1772KKKr   c                     | t        |       z  S r   )rU   r   s    r   	hard_silurW   q   s    |Ar   c                     t         j                  j                  |       }|dk(  r|S t        j                  | dkD  |||z        S )Nr.   r   )r   r   elur1   )r    rQ   ress      r   rY   rY   u   s:    
%%))A,Cz
xxAsECK00r   c                 @    t         j                  j                  |       S r   )r   r   selur   s    r   r\   r\   }   r!   r   c                 Z    t        |       } t        j                  j                  | |      S )N)approximate)r	   r   r   gelu)r    r^   s     r   r_   r_      s#    !A55::a[:11r   c                     t        j                  | d      |t         j                  j                  t        j                  | d      |z        z  z   S )Nr+   )r   maximumr7   expm1minimum)r    rQ   s     r   celurd      sC    ::a


1cU") !  r   c                     | j                   |   dz  dk7  rt        d| j                    d|       t        j                  | d|      \  }}|t        j                  |      z  S )NrC   r   z4axis size must be divisible by 2. Received: x.shape=z with axis=)num_or_size_splitsaxis)shape
ValueErrorr   splitr%   )r    rg   x1x2s       r   glurm      sh    wwt}qA!!"	TF<
 	
 XXaAD9FB

2r   c                 2    t        j                  | dd      S )Ng      r/   )clip_value_minclip_value_max)r   clip_by_valuer   s    r   	hard_tanhrr      s    Ad3GGr   c                     t        j                  t        j                  |       |kD  | t        j                  |             S r   )r   r1   absr>   r?   s     r   hard_shrinkru      s+    88BFF1I	)1bmmA.>??r   c                 6    t        j                  | |kD  | |      S r   )r   r1   )r    r@   default_values      r   r@   r@      s    88A	M1m44r   c                 "   | }|bt        j                  | dg      }t         j                  j                  |d      }t        j                  |t        j                  |             }n!t         j                  j                  | |      }||_        |S Nr*   rg   )r   reshaper   softmaxrh   r&   )r    rg   r'   r(   s       r   r|   r|      sq    F| At$vB/FBHHQK0qt,!FMr   c                 
   |at        j                  | dg      }t         j                  j                  |d      }t        j                  |t        j                  |             S t         j                  j                  | |      S ry   )r   r{   r   log_softmaxrh   )r    rg   r(   s      r   r~   r~      sg    | At$""6"3zz&"((1+..55QT**r   c                 |   t        |       }t        j                  |d|      }t        j                  ||      }t        j                  dt        j
                  |      |   dz   |j                        }dgt        |j
                        z  }d||<   t        j                  ||      }||dz
  |z  z
  dkD  }t        j                  ||d      }t        j                  t        j                  ||j                        |d	
      }	t        j                  ||d	
      dz
  |	z  }
t        j                  ||
z
  d      }|S )N
DESCENDING)	directionrg   rz   r.   r,   r*   r   r+   Trg   keepdims)r	   r   sortcumsumrangerh   r-   lenr{   r1   
reduce_sumr   ra   )r    rg   r'   logits_sortedlogits_cumsumrr_shapesupportlogits_cumsum_safektaur(   s               r   	sparsemaxr      s   q!FGGFlFMIIm$7M
BHHV$T*Q.fllCAcC%%GGDM


1gA}q0A559G'=#>
bgggv||44$OA==+$FJa
OCZZc*FMr   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 t	        d| j                   d      )	NrC   r.   r   rC   r.   r   rC      r.   r   r   rC   r      r.   zePooling inputs's shape must be 3, 4 or 5, corresponding to 1D, 2D and 3D inputs. But received shape: .)r   rh   r   	transposeri   )inputsnum_spatial_dimss     r   _transpose_spatial_inputsr      s    6<<(1, 1fi0 M 
Q	fl3 M 
Q	fo6 M	 228,,qB
 	
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 )NrC   r.   r   r   r   r.   rC   r   r   r   r.   rC   r   )r   rh   r   r   )outputsr   s     r   _transpose_spatial_outputsr      su    7==)A-1,,w	2
 N	 
Q	,,w5 N 
Q	,,w8Nr   c                     t        j                  |      }||n|}|j                         }t        dt	        | j
                              }|dk(  rt        |       } t        j                  j                  | ||||      }|dk(  rt        |      }|S Nchannels_lastchannels_first)r   standardize_data_formatupper_convert_data_formatr   rh   r   r   r   max_poolr   r   	pool_sizestridespaddingdata_formattf_data_formatr   s          r   r   r           11+>K"?iGmmoG)/3v||;LMN&& +62eennG &&,W5Nr   c                     t        j                  |      }||n|}|j                         }t        dt	        | j
                              }|dk(  rt        |       } t        j                  j                  | ||||      }|dk(  rt        |      }|S r   )r   r   r   r   r   rh   r   r   r   avg_poolr   r   s          r   average_poolr     r   r   c                 2   t        j                  t        j                  t        j                  t        j                  |      t         j                        t        j                  | t         j                        z  t        j                  |t         j                        z        t         j
                        }t        j                  t         j                  j                  t        j                  t        j                  d|dz         t         j                        t        j                  | t         j                        z  t        j                  |t         j                        z        t         j
                        }t        j                  ||       }t        j                  || dz
        }||z
  }t        j                  ||      }t        d| |z
  dz         }|}	||z   }
t        j                  ||
|	      }t        j                  |t         j
                        S )z>Compute gather indices for Two-Pool Gather method (corrected).r.   )r   r   floorr   float32int32r7   ceilrc   equalmaxr1   )	input_dimoutput_sizesmall_window
big_windowwindow_startswindow_endswindow_sizesis_big_windowsmall_pool_lensmall_indicesbig_indicesgather_indicess               r   _compute_static_gather_indicesr   "  st    GG
GGBHH[)2::6ggi,-ggk2::./	

 	M ''
GGBHHQa0"**=ggi,-ggk2::./	

 	K **[)4KJJ}i!m<M.LHH\:6MI4q89N!M.0KXXm[-HN77>288,,r   c                 .   t        |t              r|f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|t        d      t        ||      \  }}t        ||||      }t        j                  j                  | |fdddd	
      }	t        j                  j                  | |fdddd	
      }
t        j                  |	|
gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )Nr   r   r.   r   :Input length must be statically known for adaptive poolingAVGr.   VALIDNWCwindow_shapepooling_typer   r   r   rz   
isinstanceintr   r   rh   as_listri   r   r   r   poolconcatgatherr   r   r   static_shapel_staticout_lsmall_lbig_lgather_lsmall_pool_l
big_pool_l
combined_lpooled_ls                r   _adaptive_average_pool1dr   F  '   +s#"n&&fi0<<'')LAHNEH
 	
 ;8UKNGU-hwNH55::Z  L X  J L*5A>JyyXA6H&&<<)4Or   c                 .   t        |t              r|f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|t        d      t        ||      \  }}t        ||||      }t        j                  j                  | |fdddd	
      }	t        j                  j                  | |fdddd	
      }
t        j                  |	|
gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )Nr   r   r.   r   r   MAXr   r   r   r   rz   r   r   s                r   _adaptive_max_pool1dr   q  r   r   c                 z   t        |t              r||f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|\  }}||t        d      t        ||      \  }}	t        ||      \  }
}t        ||||	      }t        |||
|      }t        j                  j                  | |dfdddd	
      }t        j                  j                  | |	dfdddd	
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |d|
fdddd	
      }t        j                  j                  |d|fdddd	
      }t        j                  ||gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )Nr   r   r.   rC   FInput spatial dimensions must be statically known for adaptive poolingr   r.   r.   r   NHWCr   rz   r   r   r   r   r   r   h_staticw_staticout_hout_wsmall_hbig_hsmall_wbig_wgather_hgather_wsmall_pool_h
big_pool_h
combined_hpooled_hsmall_pool_w
big_pool_w
combined_wpooled_ws                         r   _adaptive_average_pool2dr     s   +s#"K0&&fl3<<'')LAHAHLE58+4
 	

 ;8UKNGU:8UKNGU-hwNH-hwNH55::q\  L QZ  J L*5A>JyyXA6H55::\  L Z  J L*5A>JyyXA6H&&<<,7Or   c                 z   t        |t              r||f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|\  }}||t        d      t        ||      \  }}	t        ||      \  }
}t        ||||	      }t        |||
|      }t        j                  j                  | |dfdddd	
      }t        j                  j                  | |	dfdddd	
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |d|
fdddd	
      }t        j                  j                  |d|fdddd	
      }t        j                  ||gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )z5Adaptive Max Pooling 2D using Two-Pool Gather method.r   r   r.   rC   r   r   r   r   r   r   rz   r   r   r   s                         r   _adaptive_max_pool2dr    s   +s#"K0&&fl3<<'')LAHAHLE58+4
 	

 ;8UKNGU:8UKNGU-hwNH-hwNH55::q\  L QZ  J L*5A>JyyXA6H55::\  L Z  J L*5A>JyyXA6H&&<<,7Or   c                    t        |t              r|||f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|d   }|\  }}}	|||t        d      t        ||      \  }
}t        ||      \  }}t        ||	      \  }}t        |||
|      }t        ||||      }t        ||	||      }t        j                  j                  | |
ddfddd	d
      }t        j                  j                  | |ddfddd	d
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |d|dfddd	d
      }t        j                  j                  |d|dfddd	d
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |dd|fddd	d
      }t        j                  j                  |dd|fddd	d
      }t        j                  ||gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )Nr   r   r.   rC   r   r   r   r.   r.   r.   r   NDHWCr   rz   r   r   r   r   r   r   d_staticr   r   out_dr   r   small_dbig_dr   r   r   r   gather_dr   r   small_pool_d
big_pool_d
combined_dpooled_dr   r   r   r   r   r   r   r   s                                  r   _adaptive_average_pool3dr  )  s   +s#"K=&&fo6<<'')LAHAHAH%E5%8+x/?4
 	

 ;8UKNGU:8UKNGU:8UKNGU-hwNH-hwNH-hwNH55::q!_  L Q]  J L*5A>JyyXA6H55::!_  L ]  J L*5A>JyyXA6H55::G_  L E]  J L*5A>JyyXA6H&&<</:Or   c                    t        |t              r|||f}|dk(  rt        j                  | d      } | j                  j                         }|d   }|d   }|d   }|\  }}}	|||t        d      t        ||      \  }
}t        ||      \  }}t        ||	      \  }}t        |||
|      }t        ||||      }t        ||	||      }t        j                  j                  | |
ddfddd	d
      }t        j                  j                  | |ddfddd	d
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |d|dfddd	d
      }t        j                  j                  |d|dfddd	d
      }t        j                  ||gd      }t        j                  ||d      }t        j                  j                  |dd|fddd	d
      }t        j                  j                  |dd|fddd	d
      }t        j                  ||gd      }t        j                  ||d      }|dk(  rt        j                  |d      }|S )z5Adaptive Max Pooling 3D using Two-Pool Gather method.r   r   r.   rC   r   r   r   r  r   r  r   rz   r   r   r  s                                  r   _adaptive_max_pool3dr    s   +s#"K=&&fo6<<'')LAHAHAH%E5%8+x/?4
 	

 ;8UKNGU:8UKNGU:8UKNGU-hwNH-hwNH-hwNH55::q!_  L Q]  J L*5A>JyyXA6H55::!_  L ]  J L*5A>JyyXA6H55::G_  L E]  J L*5A>JyyXA6H&&<</:Or   c                     t        j                  |      }t        | j                        dz
  }|dk(  rt	        | ||      S |dk(  rt        | ||      S |dk(  rt        | ||      S t        d      )NrC   r.   r   z9adaptive_average_pool supports 1D, 2D, or 3D inputs only.)r   r   r   rh   r   r   r  ri   r   r   r   ndimss       r   adaptive_average_poolr    sw    11+>K!Ez'[II	!'[II	!'[IIG
 	
r   c                     t        j                  |      }t        | j                        dz
  }|dk(  rt	        | ||      S |dk(  rt        | ||      S |dk(  rt        | ||      S t        d      )NrC   r.   r   z5adaptive_max_pool supports 1D, 2D, or 3D inputs only.)r   r   r   rh   r   r  r  ri   r  s       r   adaptive_max_poolr    sw    11+>K!Ez#FKEE	!#FKEE	!#FKEEC
 	
r   c                     | dk(  r!|dk(  ry|dk(  ry|dk(  ryt        d| d	      | d
k(  r!|dk(  ry|dk(  ry|dk(  ryt        d| d	      t        d|  d      )Nr   r   r   r   r      r  zInput rank not supported: z. Expected values are [3, 4, 5]r   NCWNCHWNCDHWzInvalid data_format: z9. Expected values are ["channels_first", "channels_last"])ri   )r   ndims     r   r   r     s    o%19QYQY,TF 30 0  
(	(19QYQY,TF 30 0 
 #K= 1F F
 	
r   c                 \   	  fd	t        j                  d      	fd       }dk(  xr t         j                        dk(  }t	        j
                        dk(  r j                  d   }n j                  d	   }|xs |j                  d
   k7  }|r |       S  	       S )Nc                     t        t        j                              } t        j                  j                  j                         |       }|j                  }|j                         ret        j                  |j                               dk(  r?t        j                  j                  j                               dk7  rt        d| d      |S )Nr   	dilationsr   zEThe convolution operation resulted in an empty output. Output shape: z. This can happen if the input is too small for the given kernel size, strides, dilation rate, and padding mode. Please check the input shape and convolution parameters.)r   r   rh   r   r   convolutionr   is_fully_definedr7   prodr   ri   )	r   resultresult_shaper   dilation_rater   kernelr   r   s	      r   _convzconv.<locals>._conv*  s    -k3v||;LM""MMO&# # 
 ||))+		,..01Q6		&,,..01Q6 > "  r   T)jit_compilec                               S r   r   )r)  s   r   	_conv_xlazconv.<locals>._conv_xlaF  s
    wr   r   r  r   r*   r.   )r   functionr   rh   r   r   )
r   r(  r   r   r   r'  r,  	needs_xlachannelsr)  s
   ``````   @r   convr1  "  s     8 [[T" # //JC4E4JI11+>Ko%<<#<<?9Xb)99I{wr   c                 T    |yt        d | D              xr t        d |D              S )NFc              3   &   K   | ]	  }|d kD    ywr.   Nr   )r   ss     r   r   zA_needs_depthwise_stride_dilation_decomposition.<locals>.<genexpr>_  s     &q1u&   c              3   &   K   | ]	  }|d kD    ywr4  r   r   s     r   r   zA_needs_depthwise_stride_dilation_decomposition.<locals>.<genexpr>_  s     .Lq1u.Lr6  )any)r   r'  s     r   ._needs_depthwise_stride_dilation_decompositionr9  Y  s.     &g&&L3.Lm.L+LLr   c                 t   t        j                  | t         j                        dd }t        j                  ddgt         j                        g}t	        |      D ]|  \  }}||   }||   }	||   }
|dz
  |
z  dz   }||	z   dz
  |	z  }t        j
                  |dz
  |	z  |z   |z
  d      }|dz  }||z
  }|j                  t        j                  ||g             ~ |j                  t        j                  ddgt         j                               t        j                  | t        j                  |            S )N)out_typer.   r*   r   r,   rC   )	r   rh   r   r2   	enumeratera   appendstackpad)r   kernel_sizer   r'  spatial_shapepaddingsir   nr5  r   effective_kernel_sizeout_size	total_pad
pad_before	pad_afters                   r   _pad_same_spatial_channels_lastrJ  b  s$   
 HHVbhh7"=MQF"((34H+& ;1!AJ!!"Q!aEAI!#JJ\Q!66:A
	 !^

*	*i!89:; OOBKKAbhh7866&"((8,--r   c                    t        | j                        dz
  }|j                         }|dk(  rt        |       } |dk(  rVt	        j
                  | d      } t	        j
                  |d      }|t	        j
                  |d      }d|d   f}d|d   f}	n|}|}	t        d |j                  d d D              }
|dk(  rt        | |
||	      } n|dk7  rt        d	|       t        j                  j                  | |d
dd|	      }|d d d d |d   d d |d   d d f   }|$t        j                  j                  ||d
dd      }|dk(  rt	        j                  |dg      }|dk(  rt        |      }|S )NrC   r   r.   rz   r   c              3   2   K   | ]  }t        |        y wr   )r   )r   rC  s     r   r   z@_depthwise_conv_stride_dilation_decomposition.<locals>.<genexpr>  s     C1ACs   SAMEr   z7`padding` must be 'valid' or 'same'. Received: padding=)r.   r.   r.   r.   r   )r   r   r   r!  )r   r   r   )r   rh   r   r   r   expand_dimstuplerJ  ri   r   depthwise_conv2dconv2dsqueezer   )r   depthwise_kernelr   r   r   r'  pointwise_kernelr   spatial_stridesspatial_dilation_rater@  r   s               r   -_depthwise_conv_stride_dilation_decompositionrW  {  s    6<<(1,mmoG&&*621Q/>>*:C'!~~.>QGgaj/!"M!$4 5! -C(8(>(>r(BCCK&0K2G
 
G	EgYO
 	
 ee$$' % G a.OA..0E?13E0EqHIG#%%,,   
 1**Wqc*&&,W5Nr   c                 V   t        j                  |      }t        |       } t        |      }t        | j                        dz
  }|dkD  rt        d| j                   d      t        | ||       t        |d      }|j                         }t        |t              r|f|z  }t        |t              r|f|z  }t        ||      rt        | |||||      S |dk(  r|dk(  xr
 t               }|rt        |       } |s|dk(  rd|dz  z   dz   }d}	n
d	|dz  z   }d}	t!        j"                  | |	      } t!        j"                  |d
      }|d nd|z   }|s|dk(  rt        dd      }
nt        dd      }
t         j$                  j'                  | ||||
|      }t!        j(                  ||	g      }|rt+        |      }|S |dk(  xr
 t               }|rt        |       } |s|dk(  rd|z   dz   }t        dd      }
nd	|z   }|}
t         j$                  j'                  | ||||
|      }|rt+        |      }|S )NrC   z<`inputs` rank must be 3 (1D conv) or 4 (2D conv). Received: r   r   r.   r   r   r   r   r   rz   r   )r   r   r	   r   rh   ri   r  r   r   r   r   r   r9  rW  r   r   r   rN  r   rP  rR  r   )r   r(  r   r   r   r'  r   r   need_transposespatial_start_dimconv_data_formatr   s               r   depthwise_convr\    s    11+>Kv&Fv&F6<<(1,!J{{m1
 	
 ffk: *+q9NmmoG'3*//-%&(+;;5g}M<FGWk=
 	
 1
 %(88HY[.v6F[O;Wq[(4/G !w{*G !(9:Q/ - 54-;O[O;3OQG34DaH%%(((# ) 
 **W'8&9:09G !$44DN*627.4'/C7")ee$$$ % G ,W5Nr   c           	      ~   t        j                  |      }t        |       } t        |      }t        |      }t        | j                        dz
  }|dkD  rt        d| d      t        | ||       t        |d      }|j                         }t        |t              r|f|z  }t        |t              r|f|z  }t        ||      rt        | ||||||      S |dk(  r|dk(  xr
 t               }	|	rt        |       } |	s|dk(  rd	|dz  z   d	z   }d}
t        dd      }nd
|dz  z   }d}
t        dd      }t        j                   | |
      } t        j                   |d      }t        j                   |d      }|d nd	|z   }t        j"                  j%                  | ||||||      }t        j&                  ||
g      }|	rt)        |      }|S |dk(  xr
 t               }	|	rt        |       } |	s|dk(  rd	|z   d	z   }t        dd      }nd
|z   }|}t        j"                  j%                  | ||||||      }|	rt)        |      }|S )NrC   z>`num_spatial_dims` must be 1 or 2. Received: num_spatial_dims=r   r   )rT  r.   r   r   r   r   r   rz   r   )r   r   r	   r   rh   ri   r   r   r   r   r   r9  rW  r   r   r   rN  r   separable_conv2drR  r   )r   rS  rT  r   r   r   r'  r   r   rY  rZ  r[  r   s                r   separable_convr_    s    11+>Kv&F()9:()9:6<<(1,!  014
 	
 f&6D *+q9NmmoG'3*//-%&(+;;5g}M<-
 	
 1 %(88HY[.v6F[O;Wq[(4/G !3OQGw{*G !34DaH(9:>>*:C>>*:C - 54-;O%%(((# ) 
 **W'8&9:09G
 !$44DN*627.4'/C7")ee$$$ % G ,W5Nr   c           
         t        j                  |      }t        |       } t        |      }t        | ||       t	        |t        | j                              }|j                  d d }|j                  d   }	t        | j                        }
t        j                  |       }t        |
      D ]  \  }}|	||   |
|<    t        |
||	|||||      }t        j                  j                  | ||||j                         ||      S )Nr-  )r   r   r!  )r   r   r	   r   r   r   rh   listr   r<  r   r   conv_transposer   )r   r(  r   r   output_paddingr   r'  r   r@  filtersinput_shapesymbolic_shaperC  eoutput_shapes                  r   rb  rb  }  s    11+>Kv&Fv&F'D)+s6<<7HIN,,s#Kll2Gv||$KXXf%N+& /19+A.KN/ 7	L 55"    r   c                    t        | d      } |d}nt        j                  |      }|r|dk  r|t        | j                        z   dz   }t        j                  | j                        }t        j                  | |f      }t        j                  t        j                  |d      |      }| j                  D cg c]  }t        j                  |       }}t        j                  |ddi}|j                  |t        j                  | d             |D 	cg c]  }	t        j                  |	|df       }}	|D 	cg c]&  }	t        j                  |	t        j                        ( }}	t        j                   |d      }t#        | j                        }
|
j                  ||       t        j$                  |||
      S |d	k(  rd
nd\  }}t        j&                  | |||||      S c c}w c c}	w c c}	w )Nint64r,   r   r   r.   indexingijrz   bool)TF)NN)on_value	off_valuerg   r-   )r	   r   standardize_dtyper   rh   r7   r$  r   r{   r   greater_equalr   meshgridinsertra   rj  r   ra  SparseTensorone_hot)r    num_classesrg   r-   sparsevalues_countvaluesdimindicesarh   rn  ro  s                r   ru  ru    s   !7+A}))%0 !8#agg,&*Dyy)A/ ))&!4EB,-GG4S288C=44++w66tRZZ1-.=DE2::a,!23EE189A2771bhh'99))G!,QWWT;'w66+0F?-Hi::	  5 F9s   :G%G*<+G/c                 <   t        | j                        dkD  rdnd}t        j                  |      dk(  r|rgt	        | ||dd      }t
        j                  j                  ||d      }|j                  }t        j                  ||      }|j                  |       |S t	        | |||      }t        j                  ||	      S |r2t	        | |||d      }t
        j                  j                  ||d      S t	        | |||      }t        j                  ||	      S )
Nr.   r   rm  int8T)rg   r-   rw  )rg   output_is_sparse)rg   r-   rz   )r   rh   r   rp  ru  r   rw  
reduce_maxr   	set_shape
reduce_any)r    rv  rg   r-   rw  reduction_axisr   outputs_shapes           r   	multi_hotr    s   agg,*QN  '61 ;TG ii**nt + G $MMMgggu-Gm,Na4uEG==~>> ;TtG 99''nt (   a4uEG==~>>r   c                    | }|}t        | d      }|r| j                  }d}t        | d      xrP t        | t        j                  j
                  t        j                  f       xr | j                  j                  |k(  xr | }|rLt        | j                  j                        dk7  rt        d| d      | j                  j                  d   }d}|r"|s|rt        j                  d| d	| d
d       ||fS )zCRetrieves logits tensor from maybe-softmax or maybe-sigmoid tensor.r&   Topr.   zExpected 1 input for r   r   z"`zK` received `from_logits=True`, but the `output` argument was produced by a zB activation and thus does not represent logits. Was this intended?rC   )
stacklevel)hasattrr&   r   r   __internal__EagerTensorVariabler  typer   r   ri   warningswarn)r(   from_logitsop_typefn_nameoutput_from_logits_has_keras_logitsfrom_expected_op_types           r   _get_logitsr    s   GLv7&& 	 	&6BOO$?$?#MNN	&IINNg% 
	   vyy A%4WIQ?@@))""1%(,A	 77>i @!! 	
 L  r   c                    t        j                  |       } t        j                  |      }t        | j                        dk  r%t	        d| j                   d|j                         t        | j                        t        |j                        k7  r%t	        d| j                   d|j                         t        | j                  |j                        D ]5  \  }}|	|||k7  st	        d| j                   d|j                          t        ||dd      \  }}|r"t         j                  j                  | ||      S |t        j                  ||d	
      z  }t        j                  |t        j                         dt        j                         z
        }t        j                  | t         j                  j                  |      z  |       S )a  Categorical crossentropy between an output tensor and a target tensor.

    Args:
        target: A tensor of the same shape as `output`.
        output: A tensor resulting from a softmax
            (unless `from_logits` is `True`, in which
            case `output` is expected to be the logits).
        from_logits: Boolean, whether `output` is the
            result of a softmax, or is a tensor of logits.
        axis: Int specifying the channels axis. `axis=-1` corresponds to data
            format `channels_last`, and `axis=1` corresponds to data format
            `channels_first`.

    Returns:
        Output tensor.

    Example:

    >>> a = tf.constant([1., 0., 0., 0., 1., 0., 0., 0., 1.], shape=[3,3])
    >>> print(a)
    tf.Tensor(
      [[1. 0. 0.]
       [0. 1. 0.]
       [0. 0. 1.]], shape=(3, 3), dtype=float32)
    >>> b = tf.constant([.9, .05, .05, .05, .89, .06, .05, .01, .94],
    ...                 shape=[3, 3])
    >>> print(b)
    tf.Tensor(
      [[0.9  0.05 0.05]
       [0.05 0.89 0.06]
       [0.05 0.01 0.94]], shape=(3, 3), dtype=float32)
    >>> loss = categorical_crossentropy(a, b)
    >>> print(np.around(loss, 5))
    [0.10536 0.11653 0.06188]
    >>> loss = categorical_crossentropy(a, a)
    >>> print(np.around(loss, 5))
    [0. 0. 0.]
    r.   zPArguments `target` and `output` must be at least rank 1. Received: target.shape=, output.shape=WArguments `target` and `output` must have the same rank (ndim). Received: target.shape=QArguments `target` and `output` must have the same shape. Received: target.shape=Softmaxcategorical_crossentropy)labelsr'   rg   Tr   r/   )r   r	   r   rh   ri   zipr  r   !softmax_cross_entropy_with_logitsr   rq   r   epsilonr7   log)targetr(   r  rg   e1e2s         r   r  r    s   N !!&)F!!&)F
6<<1"LL>H
 	

 6<<C--"LL>H
 	

 fllFLL1 B>bnr  &~_V\\NL  &Y(BFK uu66&t 7 
 	
 bmmFD4@@F !3):#:F MM&277;;v#66===r   c                 j   |dk7  r)|t        |j                        dz
  k7  rt        d|       t        ||dd      \  }}t	        j
                  |       } t	        j                  | d      } t	        j
                  |      }t        | j                        t        |j                        k(  r)| j                  d   dk(  rt	        j                  | d      } t        |j                        dk  rt        d	|j                         t        | j                        t        |j                  d
d       k7  r%t        d| j                   d|j                         t        | j                  |j                  d
d       D ]5  \  }}|	|||k7  st        d| j                   d|j                          |s]t	        j                  |t        j                         dt        j                         z
        }t        j                  j                  |      }t        j                  j                  | |      }|S )aN  Categorical crossentropy with integer targets.

    Args:
        target: An integer tensor.
        output: A tensor resulting from a softmax
            (unless `from_logits` is True, in which
            case `output` is expected to be the logits).
        from_logits: Boolean, whether `output` is the
            result of a softmax, or is a tensor of logits.
        axis: Int specifying the channels axis. `axis=-1` corresponds to data
            format `channels_last`, and `axis=1` corresponds to data format
            `channels_first`.

    Returns:
        Output tensor.
    r*   r.   z4Only axis=-1 is currently supported. Received: axis=r  sparse_categorical_crossentropyrj  r,   rz   zBArgument `output` must be at least rank 1. Received: output.shape=NzRArgument `output` must have rank (ndim) `target.ndim - 1`. Received: target.shape=r  zcArguments `target` and `output` must have the same shape up until the last dimension: target.shape=r  r'   )r   rh   ri   r  r   r	   r   rR  r  rq   r   r  r7   r  r   (sparse_softmax_cross_entropy_with_logits)r  r(   r  rg   r  r  r%  s          r   r  r  e  s
   " rzdc&,,/!33B4&I
 	
 &Y(IFK !!&)FWWV7+F!!&)F
6<<C--&,,r2Ba2GF,
6<<1"LL>+
 	

 6<<CSb 122"LL>H
 	

 fllFLL"$56 B>bnr  &~_V\\NL  !!GOO%q7??+<'<
 V$UU;;f < F Mr   c                 *   t        j                  |       } t        j                  |      }t        | j                        t        |j                        k7  r%t	        d| j                   d|j                         t        | j                  |j                        D ]5  \  }}|	|||k7  st	        d| j                   d|j                          t        ||dd      \  }}|r!t         j                  j                  | |      S t        j                  |t        j                         dt        j                         z
        }| t         j                  j                  |      z  }|d| z
  t         j                  j                  d|z
        z  z  }| S )	ap  Binary crossentropy between an output tensor and a target tensor.

    Args:
        target: A tensor with the same shape as `output`.
        output: A tensor.
        from_logits: Whether `output` is expected to be a logits tensor.
            By default, we consider that `output`
            encodes a probability distribution.

    Returns:
        A tensor.
    r  r  r  Sigmoidbinary_crossentropyr  r/   r.   )r   r	   r   rh   ri   r  r  r   !sigmoid_cross_entropy_with_logitsrq   r   r  r7   r  )r  r(   r  r  r  bces         r   r  r    sx    !!&)F!!&)F
6<<C--"LL>H
 	

 fllFLL1 B>bnr  &~_V\\NL  &Y(=FK uu66& 7 
 	

 !3):#:F 277;;v&
&CAJ"''++a&j111C4Kr   c                    d}t        j                  | j                        }|dv rd}t        | d      } |rt	        | ||      \  }}nt        | ||      \  }}|rt        j                  |t        j                  j                  t        j                  j                        }t        j                  |t        j                  j                  t        j                  j                        }t        ||      }t        ||      }||fS )NF)float16bfloat16Tr   )r   rp  r-   r   _compute_moments_sync_compute_momentsr   rq   r  minr   )r    axesr   synchronized	need_cast	ori_dtypemeanvariances           r   momentsr    s     I))!''2I++	I.q$Ah)!T8<hbjjnnbjjnnE##HbjjnnbjjnnMD)$),>r   c                    t         j                  j                         }|st        | ||      S t        j                  | d      }t        j
                  | |d      }t        j
                  t        j                  |       |d      }t        j
                  ||d      }|j                  t         j                  j                  j                  |      }|j                  t         j                  j                  j                  |      }|j                  t         j                  j                  j                  |      }	t         j                  j                  ||	      }
t         j                  j                  ||	      }t        j                  |t        j                  |
      z
  d      }|s,t        j                  |
|      }
t        j                  ||      }|
|fS )Ncount)nameTr   r+   )r   
distributeget_replica_contextr  	ones_liker   rJ   
all_reduceReduceOpSUMr7   divide_no_nanra   rR  )r    r  r   replica_ctxlocal_count	local_sumlocal_squared_sumy_sumy_squared_sum	count_sumr  y_squared_meanr  s                r   r  r    sY   --335K422,,qw/KadT:IbiilM--$FK ""2==#9#9#=#=yIE**
""$5M &&r}}'='='A'A;OI77  	2DWW**=)DNzz.299T?:C@Hzz$%::h->r   c                 F    t         j                  j                  | ||      S )Nr  )r   r   r  )r    r  r   s      r   r  r    s    55==D8=44r   c                 d   |dk7  rdgt        | j                        z  }|j                  d   ||<   t        j                  ||      }t        j                  ||      }|t        j                  ||      }|t        j                  ||      }t        j                  j                  | |||||      S )Nr*   r.   r   )r    r  r  offsetscalevariance_epsilon)r   rh   r   r{   r   batch_normalization)r    r  r  rg   r  r  r  rh   s           r   r  r    s     rzc!''l"jjmdzz$&::h.ZZ.FJJue,E55$$
  %  r   c                 P   t        |       } t        |      }t        j                  | d      } t        j                  |j
                  d      }|dk(  rdn|}t        j                  ||      }t        j                  j                  | ||||d      }t        j                  ||      S )Nr   r,   r   float64F)r  r'   label_lengthlogit_lengthblank_indexlogits_time_major)r	   r   r   r   result_typer-   r   ctc_loss)r  r(   target_lengthoutput_length
mask_indexresult_dtypecompute_dtypelosss           r   r  r  &  s    v&Fv&FWWV7+F &&v||Y?L!-!:IMWWV]+F55>>""  D 774&&r   c                    t        |       } t        j                  |       }|d   |d   }	}t        j                  | d      } t	        j
                  | j                  d      }
t        j                  | |
      } t        |d      }|dk(  r't        j                  j                  | |||      \  }}nx|d	k(  rd|;| d
d |f   }| d
||dz   f   }| d
|dz   d f   }t        j                  |||gd      } t        j                  j                  | |||      \  }}nt        d| d      g }|D ]_  }t        j                  |j                  |j                  ||	f      }|j!                  t        j"                  j%                  |d             a t        j&                  |d      }t        j                  |d      }|d	k(  r,|*|dk  r||d   z   }t        j(                  ||k\  |dz   |      }||fS )Nr   r.   )r.   r   rC   r   r   r,   greedy)r   sequence_lengthmerge_repeatedr  beam_search.r*   rz   )r   r  
beam_width	top_pathszInvalid strategy z2. Supported values are 'greedy' and 'beam_search'.)sp_inputrw   )r	   r   rh   r   r   r  r-   r   r   ctc_greedy_decoderr   ctc_beam_search_decoderri   rt  r{  ry  r=  rw  to_denser>  r1   )r   sequence_lengthsstrategyr  r  r  r  re  num_samples	num_stepsr-   decodedscoresinputs_beforeinputs_maskinputs_afterdecoded_densests                     r   
ctc_decoder  =  s'    v&F((6"K(^[^K\\&),Fi8EWWVU#F()9I8EE44,)"	 5 
& 
]	" !"3#34M j:>&A!ABK!#zA~'7"78LYYk:F EE99,!	 : 
& z ** *
 	
 M P__RZZ[)4LMRYY//2/NOP HH]3MGGM73M = Z%;>#k"o5JZ'):M
 &  r   c                 B   ddl m} | j                  |j                  k7  r&t        d| j                   d|j                   d      t	        ||j
                        }t        j                  t        j                  | |z
              }d ||      z  d ||      z  z
  }|S )	Nr   )log10zInput shapes z and z" must match for PSNR calculation. r,      
   )	"keras.src.backend.tensorflow.numpyr  rh   ri   r	   r-   r   reduce_meanrJ   )rk   rl   max_valr  msepsnrs         r   r  r  ~  s    8	xx288BHH:U288* 5+ +
 	

  rxx8G
..27+
,CgeCj0DKr   c                 r    t        j                  |       } | dk(  rdnd}t        j                  |dz  |       S )Nr  g    @g̓$Ggffffffr,   )r   rp  r   r2   )r-   vals     r   _get_large_negativer    s5    %%e,Ei''ZC;;sTz//r   c                    ||s| S t        j                  | d      }|t        j                  ||      }|ryt        j                  |       }|d   |d   }}t         j                  j                  t        j                  ||fd      dd      }|d d d d d d f   }t        j                  ||      }t        j                  || t        | j                              }|S )Nrm  r,   rC   r   r*   r   )
r   r  logical_andrh   linalg	band_partonesr1   r  r-   )r'   mask	is_causalcombined_masklogits_shapeTSpadded_logitss           r   _apply_masksr    s    |ILLv6M}d;xx'AQ1yy""277Aq66#:BBD$1$%}d;HHv26<<@M r   c                    t        j                  | j                  d      }t        j                  d| |d      }t        j
                  ||      }t        j                  |t        j
                  ||j                              }|4t        j                  |t        j
                  ||j                              }t        |||      }	t        j                  |	j                  d      }
t        j
                  t        j                  j                  t        j
                  |	|
      d      |j                        }t        j                  d||d      S )Nr   zBTNH,BSNH->BNTSoptimal)optimizer*   rz   zBNTS,BSNH->BTNH)r   r  r-   r   einsumr   multiplyaddr  r   r|   )querykeyvaluebiasr  r  r  logits_dtyper'   r  probs_dtypeprobss               r   _dot_product_attention_xlar     s    &&u{{I>LYY(%yIFWWV\*F[[!=>Ffll ;< y9M %%m&9&99EKGG
bggm[9CSYYE 99&uyIIr   c	           	      @   |d}|rt        d      t        |       } t        |      }t        |      }t        | j                        dk7  r3t        d| j                   d|j                   d|j                   d      t	        j
                  | j                  |j                  |j                        }	t        | |	      } t        ||	      }t        ||	      }| j                  d   }
|j                  d   }|
A|?|
|kD  r:|d	kD  r5|
|z  }t        j                  ||d
      }t        j                  ||d
      }|t        ||	      }t        j                  |      d   }|,dt        j                  t        j                  |d            z  n|}t        | ||||||      S )NFz7Flash attention is not supported in tensorflow backend.r   zG`dot_product_attention` only supports 4D inputs. Received: query.shape=z, key.shape=z, value.shape=r   rC   r.   )repeatsrg   r,   r*   r/   r   )ri   r	   r   rh   r   r  r-   r   r   repeatrI   r   )r  r  r  r  r  r  r  flash_attentionattn_logits_soft_capr  num_query_headsnum_kv_headsgroupsHs                 r   dot_product_attentionr*    s    E
 	
 e$E
C
 Ce$E
5;;1%%*[[Mcii[ I ;;-q*
 	

 ''SYYLM&E
sM
"C&Ekk!nO99Q<L#$l*1 L0iiV!4		%a8 ];
bA6;mS2772771i011E%sE4y% r   c           	         t        |t              r||fn|}t        |t              r||fn|}t        |t              r||fn|}t        |t              r||fn|}| j                  \  }	}
}}t        d |D              r.t	        j
                  | ddgddg|d   |d   g|d   |d   gg      } t	        j                  | g d      }t        j                  j                  |d|d   |d   dgd|d   |d   dgd|d   |d   dgd      }|j                  \  }	}}}t	        j                  ||	|||d   |d   |
g      }t	        j                  |g d      }t	        j                  ||	|
|d   z  |d   z  ||z  g      }|S )a  Tensorflow 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)
    c              3   &   K   | ]	  }|d kD    yw)r   Nr   )r   _s     r   r   zunfold.<locals>.<genexpr>  s     
Q1q5
r6  r   r.   r   r   )imagessizesr   ratesr   )r   r  r   r   r.   rC   )
r   r   rh   r8  r   r?  r   imageextract_patchesr{   )inputr@  dilationr   strider   r   pr5  NCr)  Wr    patchesnHnWDs                     r   unfoldr>    s     k3' 
k" 
 !+8S 98xA(#6'GA&vs3AJAq!Q 
!
u1v1v!ad|adAaD\JK
UL)Ahh&&!A$!a AaD!A$"!A$!a  ' G ==LAr2qjj!RQqT1Q4+G ll#G jj1a!A$h1orBw"?@GNr   c           
          t        |t              r||fn|}t        |t              r||fn|}t        |t              r||fn|t        |t              r||fn|}t        |t              r||fn|t        j                         d   }	 j                  d   }
|\  |\  }}|
z  z  |d|d   z  z   d   dz
  z  z
  dz
  d   z  dz   |d|d   z  z   d   dz
  z  z
  dz
  d   z  dz   t        j                   |	g       |d|d   z  z   |d|d   z  z    f
d}t        j
                  |       }|d   dkD  s|d   dkD  r#|dddd|d   |d   z
  |d   |d   z
  f   }|S )a  TensorFlow implementation of Fold (col2im).
    Combine an array of sliding local blocks into a large tensor.

    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.   rC   c           	        
 t        j                  gj                        }t              D ]<  }t              D ]*  }|d   z  }|d   z  }t        j                        d   z  |z   }t        j                        d   z  |z   }| d d ||d d d d f   }t        j                  t        j                        z        }	t        j
                  t        j                  |      g      }
t        j
                  t        j
                  |g      g      }t        j                  |	|
|gd      }t        j                  |dg      }t        j                  |||      }- ? |S )Nr,   r   r.   rz   r*   )	r   zerosr-   r   r#  tiler>  r{   tensor_scatter_nd_add)x_singler(   rC  jh_startw_start	h_indices	w_indicespatchc_idxh_idxw_idxr{  ry  r8  r   kHkWr;  r<  oH_padoW_padr5  r    s                 r   _fold_singlezfold.<locals>._fold_singleS  sM   1ff-QWW=r 	KA2Y Kad(ad(HHRL1Q4/'9	HHRL1Q4/'9	 Aq!Q/		"((1+rBw7		)R 81#>	B4 81#>((E5) EB4011&'6JK	K" r   N)r   r   r   rh   r{   vectorized_map)r    r   r@  r4  r   r5  r   or6  r7  CKKoHoWrR  r(   r8  r   rN  rO  r;  r<  rP  rQ  r5  s   `              @@@@@@@@@r   foldrX  #  s   " k3' 
k"  k3' 
k" 
 !+8S 98xA(#6'GA&vs3A
AA
''!*CFBFBRA q1Q4x-!A$"q&/
)A
-!A$	6	:B
q1Q4x-!A$"q&/
)A
-!A$	6	:B 	

1q!RR,-A !ad(]F!ad(]F , |Q/F 	tax1Q4!81adVad]2AaD6AaD=4HHIMr   c                     |dk(  r"t         j                  j                  | |d      S t        j                  | g d      } t         j                  j                  | |d      } t        j                  | g d      } | S )a"  TensorFlow implementation of depth_to_space.

    Rearranges data from depth into blocks of spatial data.

    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   depth_to_spacer   r    
block_sizer   s      r   r[  r[  r  j      o%uu##Azv#FF LLL)EE  JF CLLL)r   c                     |dk(  r"t         j                  j                  | |d      S t        j                  | g d      } t         j                  j                  | |d      } t        j                  | g d      } | S )a  TensorFlow implementation of space_to_depth.

    Rearranges blocks of spatial data into depth.

    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   rZ  r   r   )r   r   space_to_depthr   r\  s      r   r`  r`    r^  r   )r0   )r   )g?)r/   )T)r*   )NvalidN)r   r   )r.   ra  Nr.   )r.   ra  NNr.   )r*   NF)Fr*   )F)FF)NNgMbP?)r   )r  d   r.   Tr   )NNNFNN)r.   r   r.   )r   )V	functoolsr7   r  
tensorflowr   	keras.srcr   &keras.src.backend.common.backend_utilsr   r   r   r   !keras.src.backend.tensorflow.corer   r	   	lru_cacher   r   r#   r%   r3   r5   r8   r:   r<   rA   rE   rG   rM   rO   rR   rU   rW   rY   r\   r_   rd   rm   rr   ru   r@   r|   r~   r   r   r   r   r   r   r   r   r   r  r  r  r  r  r   r1  r9  rJ  rW  r\  r_  rb  ru  r  r  r  r  r  r  r  r  r  r  r  r  r  r  r   r*  r>  rX  r[  r`  r   r   r   <module>ri     sR        L 3 ? T"Q #Q"5L
12
H@5+$$	 > 8!-H(V(VCLDNZz[|


F 4nM.@ AN Y@ _J (V!H?D!!HP>f<~-`.@5
 ?C.'4 >!B0*J. 
	
6r,^L^8r   