
    ijm                         d dl Z d dlmZ d dlmZ 	 	 	 	 	 	 	 	 ddZd Zd Zd Zd Z	d	 Z
d
 Z	 ddZ	 	 	 	 ddZd Zd Z	 	 	 	 ddZd Zd Zd Z	 	 ddZ	 	 	 ddZy)    N)tree)convert_to_tensorc           	         	
./01234567 xs j                   d   d }|st        j                  |      t        j                        }|d   j                   d   }|}|t|j                  t
        j                  k7  r|j                  t
        j                        }t        |j                         dk(  rt        j                  |d      }|s ||      }g dd.|r|st        d      t              }g }g }fd}t        j                        rt        j                  |      6n	 |      f66fd	}|t        j                  |      }r!t        j                  ||j                   
      }t        |      D ]  } ||      }||   3  |t        |      t              z         \  }} .3|      }|st        j                   |      }n|d   }t        j"                  |||      }t        j                  |      }t        j                  |      }t        .3fd|D              }t        d t%        |||      D              }t        j&                  ||      }
r$|j)                  |       |j)                  |       |g}|g} |d   }|d   }t        j*                  |      }	rt        j"                   .|d   |      |t        j                   |            }t        j"                   .||d      |t        j                   |            }nt        |      D ]V  } ||      }  |t        |      t              z         \  }}
r#|j)                  |       |j)                  |       Q|g}|g}X |d   }|d   }t        j*                  |      }nt              }t        fd|D              2t        j&                  |D cg c]  }|d   	 c}      }   | t              t              z         \  }!}"
r|nd}#g }$t        j                  |!      D ]M  }%t-        |%      }&t        |%      |#k  r!|&j/                  g g|#t        |%      z
  z         |$j)                  |&       O t        j0                  dt
        j2                        }'|}(n/t5        d      r!t7              t        j8                        }(n}(|Brt        j                  |dg      }t-        t        j                  |            44fd5.fd0nt;        t
        j<                        rjr_t        j8                  d      })t;        |)t
        j>                  j8                        r|)d   })t        j@                  |)dz
        77fd5nfd5d 0nd 55nt        d t        j                  |!      D              10125
 	f	d}*d}+|$|1}}},|'|k  rp|+|(k  rk |*|'|,|g| }-|-d d \  }'},}|-dd  }|+dz  }+|'|k  rH|+|(k  r)nA2
 fd}*d}+|$},|}|'|k  r,|+|(k  r' |*|'|,g| }-|-d d \  }'},|-dd  }|+dz  }+|'|k  r|+|(k  r'd /-d   }$t        /fd|$D              }t        d |D              }t        j&                  |!|      }t        j&                  |!|      }|st        j                  ||      }|||fS c c}w )N   c                     t        t        t        | j                                    }d\  |d<   |d<   t	        j
                  | |      S )N)r   r   r   r   )listrangelenshapetorchpermute)input_taxess     p/var/www/html/emotional.easysim.app/public_html/venv/lib/python3.12/site-packages/keras/src/backend/torch/rnn.pyswap_batch_timestepz rnn.<locals>.swap_batch_timestep   s=    E#gmm,-.Qa}}Wd++    r      c                    t        j                  |       rt        d|        t        j                  |      rt        d|       t        |j                        t        | j                        z
  }t        |      D ]  }t        j                  | d      }  dg|z  t        |j                  |d        z   }t        j                  | |      S )Nz:mask_t is expected to be tensor,                  but got z;input_t is expected to be tensor,                  but got r   r   )
r   	is_nested
ValueErrorr
   r   r	   r   	unsqueezer   tile)mask_tr   	fixed_dim	rank_diff_	multipless         r   _expand_maskzrnn.<locals>._expand_mask.   s    >>&!!($  >>'"")%  &V\\)::	y! 	1A__VR0F	1C)Od7==+D&EE	zz&),,r   z/Unrolling requires a fixed number of timesteps.c                 F    t        j                  |       } r| d d d   } | S )Nr   )r   unbind)r   go_backwardss    r   _process_single_input_tz$rnn.<locals>._process_single_input_tK   s&    ll7+G!$B$-Nr   c                 ^    D cg c]  }||    	 }}t        j                  |      S c c}w N)r   pack_sequence_as)timet_inpinputsprocessed_inputs      r   _get_input_tensorzrnn.<locals>._get_input_tensorX   s2    &562d86C6((55 7s   *dimsc              3   0   K   | ]  } |        y wr%    ).0sr   r   s     r   	<genexpr>zrnn.<locals>.<genexpr>r   s      %01L+%s   c              3   R   K   | ]  \  }}}t        j                  |||       ! y wr%   r   where)r1   mr2   pss       r   r3   zrnn.<locals>.<genexpr>u   s+      * 1b KK1b)*   %'r   c           	   3      K   | ]W  }st        t        j                  |            n2t        t        j                  t        j                  |d g                   Y yw)r   N)r   r   r!   flip)r1   input_r"   s     r   r3   zrnn.<locals>.<genexpr>   sR      
  $ U\\&)*%,,uzz&1#'>?@A
s   AA dtype__len__c                     |    S r%   r0   )r'   mask_tas    r   
masking_fnzrnn.<locals>.masking_fn   s    t}$r   c                 j     t         fd|D              }t        d t        |||      D              S )Nc              3   Z   K   | ]"  } |t        j                                $ yw)r:   N)r
   r   )r1   or   r   s     r   r3   z5rnn.<locals>.compute_masked_output.<locals>.<genexpr>   s.      % !c&,,6GHH%s   (+c              3   R   K   | ]  \  }}}t        j                  |||       ! y wr%   r5   )r1   r7   rF   fms       r   r3   z5rnn.<locals>.compute_masked_output.<locals>.<genexpr>   s+       1b KK1b)r9   tuplezip)r   flat_out	flat_masktiled_mask_tr   s   `   r   compute_masked_outputz"rnn.<locals>.compute_masked_output   s?    $ %%%    $'h	$J  r   dimc                 0    t        j                  |       S r%   )r   less)r'   rev_input_lengths    r   rC   zrnn.<locals>.masking_fn   s     ::&6==r   c                 0    t        j                  |       S r%   )r   greater)r'   input_lengths    r   rC   zrnn.<locals>.masking_fn   s     ==t<<r   c                 @     t         fdt        ||      D              S )Nc              3   R   K   | ]  \  }}t        j                  ||         y wr%   r5   )r1   rF   zor   s      r   r3   z5rnn.<locals>.compute_masked_output.<locals>.<genexpr>   s*      B KK2.s   $'rI   )r   rL   rM   s   `  r   rO   z"rnn.<locals>.compute_masked_output   s$     #&x#;  r   c              3   F   K   | ]  }t        j                  |        y wr%   )r   
zeros_liker1   rF   s     r   r3   zrnn.<locals>.<genexpr>   s       %()  #%s   !c                 (  	  t         fdD              }t        j                  |      }        } |t        |      t              z         \  }}t        j                  |      }rnt        j                  |      }	 |||	      }
t        j                  |      }t        j                  |      } |||      }t        j                  ||      }r nd}t	        ||
      D ]
  \  }}|||<     dz   |t        |
      ft        |      z   S )as  RNN step function.

                Args:
                    time: Current timestep value.
                    output_ta_t: TensorArray.
                    prev_output: tuple of outputs from time - 1.
                    *states: List of states.

                Returns:
                    Tuple: `(time + 1, output_ta_t, output) + tuple(new_states)`
                c              3   (   K   | ]	  }|     y wr%   r0   r1   tar'   s     r   r3   z%rnn.<locals>._step.<locals>.<genexpr>       %B2bh%B   r   r   rJ   r   r&   flattenrK   )r'   output_ta_tprev_outputstatescurrent_inputr   output
new_statesflat_outputflat_mask_outputflat_new_output
flat_stateflat_new_stateflat_final_stateta_index_to_writera   outrO   	constantsflat_zero_outputinput_tar*   rC   return_all_outputsstep_functionzero_output_for_masks   `                r   _stepzrnn.<locals>._step  s0    !&%B%B B $ 5 5fm L#D)%2!5=53C#C&"
 #ll62 , %k2 !
 #8K)9#
 "\\&1
!%j!9#8NJ$  "22:?OP
,>DA!";@ 0GB,/B()0 q+u_/EFJ  r      c                     t         fdD              }t        j                  |      } |t        |      t              z         \  }}t        j                  |      }t        j                  |      }r nd}t	        ||      D ]
  \  }	}
|
|	|<    t        j                  |      } dz   |ft        |      z   S )a)  RNN step function.

                Args:
                    time: Current timestep value.
                    output_ta_t: TensorArray.
                    *states: List of states.

                Returns:
                    Tuple: `(time + 1,output_ta_t) + tuple(new_states)`
                c              3   (   K   | ]	  }|     y wr%   r0   r`   s     r   r3   z%rnn.<locals>._step.<locals>.<genexpr>J  rb   rc   r   r   rd   )r'   rf   rh   ri   rj   rk   rp   rl   rr   ra   rs   rt   initial_statesrv   r*   rw   rx   s   `          r   rz   zrnn.<locals>._step?  s     !&%B%B B $ 5 5fm L%2!5=53C#C&"
 "&j!9"ll62,>DA!";< 0GB,/B()0 "22"N
 q+.z1BBBr   c                     t        | D cg c]  }|j                   c}      }g }t        |       D ]&  \  }}|j                  |k(  s|j                  |       ( t	        j
                  |      S c c}w r%   )maxndim	enumerateappendr   stack)tensor_listt	max_ndimsmax_listis        r   _stackzrnn.<locals>._stackd  sf    [9QVV9:IH!+. '166Y&OOA&' ;;x(( :s   A/c              3   .   K   | ]  } |        y wr%   r0   )r1   rF   r   s     r   r3   zrnn.<locals>.<genexpr>n  s     5aq	5s   c              3   &   K   | ]	  }|d      yw)r   Nr0   r]   s     r   r3   zrnn.<locals>.<genexpr>o  s     3aAbE3s   )r   )!r   r   map_structurere   r?   r   booltyper
   r   r   rJ   r   r!   r<   r	   r\   r6   rK   r&   r   r   r   extendtensorint32hasattrr   r   
isinstanceTensorreturn_typessubtract)8rx   r*   r~   r"   maskrt   unrollrW   
time_majorry   rw   r   flattened_inputs
time_stepstime_steps_trh   successive_statessuccessive_outputsr#   r,   	mask_listr   r)   rj   rk   rN   rg   flat_statesflat_new_statesflat_final_stateslast_outputoutputsinput_time_zerooutput_time_zeror   output_ta_size	output_tars   out_listr'   max_iterationsmax_lenrz   itrf   final_outputsr   r   rO   ru   rv   r   rB   rC   r+   rT   s8   ```` ` ` ``                                   @@@@@@@@@@r   rnnr      s     26<<?L, ##$7@||F+!!$**1-JL::#99UZZ(Dtzz?a??4,D&t,D	-" NOO~&	 >>&!"00'O  7v>@O	6 T*I!JJyyG	:& !1'*"1%2vy)99&"
  ,FF;)"'"2"26":K"4R"8K\6;G"ll62"&,,z":$ %5@%   %* *$'$o{%* %! ..v7HI%&--f5%,,V4*0&)/%C!1D -R0K*2.Jkk"45G##kk 2<$$[1
  ++ w!<$$W- :& 
1'*!.vy)99" &&--f5%,,V4*0&)/%
1 -R0K*2.Jkk"45G ~&  
 +
 
 //'78SV8
 ,U>2U95EE
! *<	<< 01 	'CCyH3x.(S(A BCX&		' ||AU[[1)N|Y/0>!&<!8!-zz$,5<<-.G% ell3))La8gu'9'9'='=>%ajG#(>>'A+|#L >
= J!  % %-1\\:J-K%  , ,\ B  &1K
 %"~*= %+{!5?! 2?r1B.k;*12.
a %"~*=C C8 B#KJ%"~*= %dK E* E$1"1$5!k*12.
a	 %"~*=	) "!$	59553733''(8'B++,<kJ$$%8'B++E 9s    [
c                 &   | j                   d   }t        j                  | d      }| j                   d   }t        j                  || j                        j                  |d      }||j                  d      k  }t        j                  | |k(        S )aJ  Check the mask tensor and see if it right padded.

    cuDNN uses the sequence length param to skip the tailing
    timestep. If the data is left padded, or not a strict right padding (has
    masked value in the middle of the sequence), then cuDNN won't work
    properly in those cases.

    Left padded data: [[False, False, True, True, True]].
    Right padded data: [[True, True, True, False, False]].
    Mixture of mask/unmasked data: [[True, False, True, False, False]].

    Note that for the mixed data example above, the actually data RNN should see
    are those 2 Trues (index 0 and 2), the index 1 False should be ignored and
    not pollute the internal states.

    Args:
        mask: the Boolean tensor with shape [batch, timestep]

    Returns:
        boolean scalar tensor, whether the mask is strictly right padded.
    r   rP   r   )device)r   r   sumaranger   repeatr   all)r   max_seq_lengthcount_of_true
batch_sizeindicesright_padded_masks         r   _is_sequence_right_paddedr   z  s    . ZZ]NIId*MAJll>$++>EEAG  -"9"9!"<<99T..//r   c                 X    t        j                  t        j                  |  d            S )aX  Check if input sequence contains any fully masked data.

    cuDNN kernel will error out if the input sequence contains any fully masked
    data. We work around this issue by rerouting the computation to the
    standard kernel until the issue on the cuDNN side has been fixed. For a
    fully masked sequence, it will contain all `False` values. To make it easy
    to check, we invert the boolean and check if any of the sequences has all
    `True` values.

    Args:
        mask: The mask tensor.

    Returns:
        A boolean tensor, `True` if the mask contains a fully masked sequence.
    r   rP   )r   anyr   )r   s    r   _has_fully_masked_sequencer     s       99UYYu!,--r   c                 v    t        |        }t        |       }||z  }|j                         sd}t        |      y )Na  You are passing a RNN mask that does not correspond to right-padded sequences, while using cuDNN, which is not supported. With cuDNN, RNN masks can only be used for right-padding, e.g. `[[True, True, False, False]]` would be a valid mask, but any mask that isn't just contiguous `True`'s on the left and contiguous `False`'s on the right would be invalid. You can pass `use_cudnn=False` to your RNN layer to stop using cuDNN (this may be slower).)r   r   itemr   )r   no_fully_maskedis_right_paddedvaliderror_messages        r   _assert_valid_maskr     sI    1$77O/5Oo-E::<B 	 '' r   c                 X    |sdnd}t        j                  | j                         |      S )aa  Calculate the sequence length tensor (1-D) based on the masking tensor.

    The masking tensor is a 2D boolean tensor with shape [batch, timestep]. For
    any timestep that should be masked, the corresponding field will be False.
    Consider the following example:
        a = [[True, True, False, False]
             [True, True, True, False]]
    It is a (2, 4) tensor, and the corresponding sequence length result should
    be 1D tensor with value [2, 3]. Note that the masking tensor must be right
    padded that could be checked by, e.g., `is_sequence_right_padded()`.

    Args:
        mask: Boolean tensor with shape [batch, timestep] or [timestep, batch]
            if time_major=True.
        time_major: Boolean, which indicates whether the mask is time major or
            batch major.

    Returns:
        sequence_length: 1D int32 tensor.
    r   r   rP   )r   r   int)r   batch_firsttimestep_indexs      r   "_compute_sequence_length_from_maskr     s$    * *QqN99TXXZ^44r   c                    | j                   j                         j                  |      }|j                   j                         j                  |      }|j                  d   }|Nt	        |      j                         j                  |      }t        j                  d|z  |j                  |      }nJt        j                  d|z  | j                  |      }t        j                  d|z  | j                  |      }||||gS )a  Prepares Keras LSTM weights for PyTorch's functional LSTM.

    Transposes weight matrices from Keras (input_dim, 4*units) to PyTorch
    (4*units, input_dim) format and returns weight tensors that maintain
    gradient connections.

    Args:
        kernel: The kernel weights tensor with shape (input_dim, 4*units).
        recurrent_kernel: The recurrent kernel weights tensor
            with shape (units, 4*units).
        bias: The bias tensor with shape (4*units,).
        device: The device to place the tensors on.

    Returns:
        A list of weight tensors [weight_ih, weight_hh, bias_ih, bias_hh]
        suitable for torch._VF.lstm.
    r      r?   r   )T
contiguoustor   r   r   zerosr?   )	kernelrecurrent_kernelbiasr   	weight_ih	weight_hhhidden_sizebias_ihbias_hhs	            r   prepare_lstm_paramsr     s    ( ##%((0I ""--/226:I"((+K#D)44699&A++O7==
 ++O6<<
 ++O6<<
 y'733r   c                      t         j                  j                         xr( t         j                  j                  j                         S r%   )r   cudais_availablebackendscudnnr0   r   r   _is_cuda_cudnn_availabler     s-    ::""$L)=)=)J)J)LLr   c                     ddl m} ddl m} | |j                  t        j                  |j                  fv xr> ||j
                  t        j
                  |j
                  fv xr | xr |xr
 t               S )Nr   )activations)ops)	keras.srcr   r   tanhr   sigmoidr   )
activationrecurrent_activationr   use_biasr   r   s         r   cudnn_okr   
  sv     & 	{''SXX>> 	' <=	' J	' 		'
 %&r   c                 @   |t         t        |      }t        |      }|t        |      }|j                  }t        |       j                  |      } t        |      j                  |      }t        |      j                  |      }|
r|rdnd}t	        j
                  | |g      } | j                  }|j                  dk(  xro t        j                  j                          xrN t        t        j                  d      xr t        j                  j                          xr t        ||||d u      }|rI| }|s| j                  ddd      }	 t        |||||||	|      \  }}}|s|j                  ddd      }|||fS t#        | ||||||||	|
      S # t         $ r Y w xY w)	Nr   r   r-   r   is_compilingr   r   return_sequencesr   )NotImplementedErrorr   r?   r   r   r<   r   r   jit
is_tracingr   compilerr   r   r   _cudnn_lstm	Exception_fallback_lstm)r*   initial_state_hinitial_state_cr   r   r   r   r   r   r   r"   r   r   compute_dtypeseq_dimr   cudnn_supportedcudnn_inputsr   r   rh   s                        r   lstmr     s   " !! v&F()9: & LLMv&))-8F'8;;MJO'8;;MJO "!F'3 ]]Fv 	
		$$&&	
 ENNN3 .++-
	
  %	
  !>>!Q2L	+6 !1	,(K& !//!Q2//    		s   0F 	FFc                 ,   |j                         dk(  r#|j                  d      }|j                  d      }nK|j                         dk(  r8|j                  d   dk(  r&|j                  ddd      }|j                  ddd      }t	        ||||      }t
        j                  j                  | ||f||d uddt        j                         dd	      \  }	}
}|
j                  d      }
|j                  d      }|	d d df   }|s|j                  d      }	||	|
|gfS 	Nr   r   r{   r           FTr   )
rQ   r   r   r   r   r   _VFr   is_grad_enabledsqueeze)r*   r   r   r   r   r   r   r   paramsr   h_nc_nr   s                r   r   r   z  s'    !)33A6)33A6				!	#(=(=a(@A(E)11!Q:)11!Q: )94HF 			/*D	
GS# ++a.C
++a.C!R%.K''*#s++r   c
                 @   |	r| j                  ddd      } t        j                  | |      }
||
|z   }
| j                  d   }|}|}g }t	        |      D ]~  }|
|   t        j                  ||      z   }t        j
                  |dd      \  }}}} ||      |z   ||       ||      z  z   } ||       ||      z  }|}|}|j                  |        t        j                  |d      }|	r|j                  ddd      }|}|s|j                  |	rdnd      }||||gfS )a  Pure-torch LSTM with pre-computed input projections.

    Used when cuDNN is not available (CPU, non-standard activations, etc.).
    Pre-computes all input projections in a single matmul across timesteps,
    then only computes the recurrent projection per step.
    r   r   r   r   rP   	r   r   matmulr   r	   chunkr   r   r   )r*   r   r   r   r   r   r   r   r   r   x_projr   hcr   r   zz_iz_fz_cz_onew_cnew_hr   s                           r   r   r     sJ   $ 1a( \\&&)F$aJAAG: 1IQ(899"[[A15S#s$S)A-0D1
sO1  %S)Ju,==q kk'q)G//!Q*K''[a@!Q''r   c           
         |r|t         t        |||
|d u      }t        | t        j                        r| j
                  j                  nd}|dk(  xr\ t        j                  j                          xr; t        t        j                  d      xr t        j                  j                          }|xr |}t        |      }t        |      }|t        |      }|j                  }t        |       j                  |      } t        |      j                  |      }|	rt        j                  | dg      } |r	 t!        | |||||| j
                        S t%        | |||||||      S # t"        $ r Y w xY w)Nr   cpur   r   r   r-   r   )r   r   r   r   r   r   r   r   r   r   r   r   r   r?   r   r<   
_cudnn_grur   _fallback_gru)r*   initial_stater   r   r   r   r   r   r   r"   r   reset_aftercudnn_activation_okinputs_device_typecudnn_runtime_okr   r   s                    r   grur    s   " $*!! #T!	 )>E  	f$ 	
		$$&&	
 ENNN3 .++-
  *>.>O v&F()9: & LLMv&))-8F%m477FM F!-	 !1}}  	 	  		s   7E% %	E10E1c                    t        j                  | dd      \  }}}t        j                  |dd      \  }}}	t        j                  |||gd      j                  j	                         j                  |      }
t        j                  |||	gd      j                  j	                         j                  |      }|t        j                  |d   d      \  }}}t        j                  |d   d      \  }}}t        j                  |||g      j	                         j                  |      }t        j                  |||g      j	                         j                  |      }nY|j                  d   }t        j                  d|z  | j                  |      }t        j                  d|z  | j                  |      }|
|||gS )a~  Prepares Keras GRU weights for PyTorch's functional GRU.

    Reorders gates from Keras [z, r, h] to PyTorch [r, z, h] format
    and returns weight tensors that maintain gradient connections.

    Args:
        kernel: The kernel weights tensor with shape (input_dim, 3*units).
        recurrent_kernel: The recurrent kernel weights tensor
            with shape (units, 3*units).
        bias: The bias tensor with shape (2, 3*units) for reset_after=True.
        device: The device to place the tensors on.

    Returns:
        A list of weight tensors [weight_ih, weight_hh, bias_ih, bias_hh]
        suitable for torch._VF.gru.
    r{   r   rP   r   r   )	r   r
  catr   r   r   r   r   r?   )r   r   r   r   z_kr_kh_kz_rr_rh_rr   r   z_bir_bih_biz_bhr_bhh_bhr   r   r   s                        r   prepare_gru_paramsr-  6  s   $ KKq1MCcKK 0!;MCc 		3S/q133>>@CCFKI		3S/q133>>@CCFKI !;;tAw2dD ;;tAw2dD ))T4./::<??G))T4./::<??G&,,Q/++O6<<
 ++O6<<
 y'733r   c                    |j                         dk(  r|j                  d      }n8|j                         dk(  r%|j                  d   dk(  r|j                  ddd      }t	        ||||      }t
        j                  j                  | |||d uddt        j                         dd	      \  }}	|	j                  d      }	|d d df   }
|s|
j                  d      }|
||	gfS r   )
rQ   r   r   r   r-  r   r  r  r  r  )r*   r  r   r   r   r   r   r  r   r  r   s              r   r  r  d  s     a%//2					!m&9&9!&<&A%--aA6(8$GF 99==D	
LGS ++a.C!R%.K''*#&&r   c                 t   | j                  ddd      } t        j                  | |      }|||d   z   }| j                  d   }	|}
g }t	        |	      D ]  }t        j                  |
|      }|||d   z   }t        j
                  ||   dd      \  }}}t        j
                  |dd      \  }}} |||z         } |||z         } ||||z  z         }||
z  d|z
  |z  z   }
|j                  |
        t        j                  |d      }|j                  ddd      }|
}|s|j                  d      }|||
gfS )a  Pure-torch GRU (reset_after=True) with pre-computed input projections.

    Used when cuDNN is not available (CPU, non-standard activations, etc.).
    Pre-computes all input projections in a single matmul across timesteps,
    then only computes the recurrent projection per step.
    r   r   r   r{   rP   g      ?r  )r*   r  r   r   r   r   r   r   r  r   r  r   r   h_projx_zx_rx_hh_zr&  h_hr  rhhr   s                           r   r  r    s[   " ^^Aq!$F
 \\&&)F$q'!aJAG: a!12d1g%FF1Iqa8S#FA15S# s+ s+a#g&ES1WN"q kk'q)GooaA&GK''*!$$r   c                    |t         t        ||||duxr |du      st         t        |      }t        |      }t        |	      }	t        |
      }
|j                  }t        | |      } t        ||      }t        ||      }t        ||      }t        ||      }| j                  }|j
                  dk7  sVt        j                  j                         s8t        t        j                  d      r$t        j                  j                         rt         t        ||||      }t        |	|
||      }||z   }t        j                  ||gd      }t        j                  ||gd      }	 t        j                  j                  | ||f|dd	d
t        j                          dd	      \  }}}|j(                  d   }|dd|f   } |d|df   }!|d   |d	   }#}"|d   |d	   }%}$| dddf   }&|!dddf   }'|r| }(|!})n"|&j+                  d	      }(|'j+                  d	      })|&|(|"|$gf|'|)|#|%gffS # t"        t$        t&        f$ r}t        d|       |d}~ww xY w)a  Fused bidirectional cuDNN LSTM for the torch backend.

    Runs forward and backward passes in a single
    ``torch._VF.lstm(..., bidirectional=True)`` call instead of dispatching
    two unidirectional LSTM calls. Backward outputs are returned in original
    time order, ready for the caller's ``merge_mode`` to consume directly.

    Args:
        inputs: Input tensor of shape ``(batch, time, features)``.
        fwd_initial_state_h: Initial hidden state for the forward direction,
            shape ``(batch, hidden)``.
        fwd_initial_state_c: Initial cell state for the forward direction,
            shape ``(batch, hidden)``.
        bwd_initial_state_h: Initial hidden state for the backward direction,
            shape ``(batch, hidden)``.
        bwd_initial_state_c: Initial cell state for the backward direction,
            shape ``(batch, hidden)``.
        mask: Sequence mask. Only ``None`` is supported; otherwise
            ``NotImplementedError`` is raised so the caller can fall back to
            the two-pass path.
        fwd_kernel: Forward input kernel, shape ``(features, 4 * hidden)``.
        fwd_recurrent_kernel: Forward recurrent kernel, shape
            ``(hidden, 4 * hidden)``.
        fwd_bias: Forward bias, shape ``(4 * hidden,)`` or ``None``.
        bwd_kernel: Backward input kernel, shape ``(features, 4 * hidden)``.
        bwd_recurrent_kernel: Backward recurrent kernel, shape
            ``(hidden, 4 * hidden)``.
        bwd_bias: Backward bias, shape ``(4 * hidden,)`` or ``None``.
        activation: Output activation. Only ``tanh`` engages cuDNN.
        recurrent_activation: Gate activation. Only ``sigmoid`` engages
            cuDNN.
        return_sequences: If ``True``, return outputs at every timestep;
            otherwise only the last timestep.
        unroll: Not supported; cuDNN requires the rolled path.

    Returns:
        A pair ``((fwd_last, fwd_outputs, [fwd_h_n, fwd_c_n]),
        (bwd_last, bwd_outputs, [bwd_h_n, bwd_c_n]))`` matching the JAX
        equivalent's return shape.
    Nr   r>   r   r   r   rP   Tr   r   z!cuDNN bidirectional LSTM failed: .r   )r   r   r   r?   r   r   r   r   r   r   r   r   r   r   r  r   r  RuntimeError	TypeErrorr   r   r   )*r*   fwd_initial_state_hfwd_initial_state_cbwd_initial_state_hbwd_initial_state_cr   
fwd_kernelfwd_recurrent_kernelfwd_bias
bwd_kernelbwd_recurrent_kernelbwd_biasr   r   r   r   r   fwd_h0fwd_c0bwd_h0bwd_c0r   
fwd_params
bwd_paramsr  h_0c_0r   r  r  er   y_fwdy_bwdfwd_h_nbwd_h_nfwd_c_nbwd_c_nfwd_lastbwd_lastfwd_outputsbwd_outputss*                                             r   bidirectional_lstmrX    s   t !!%>($*>	 "!":.J,-AB":.J,-AB$$Mv];F2-HF2-HF2-HF2-HF ]]Fv99!ENNN3++- "!$((FJ %((FJ *$F ++vv&A
.C
++vv&A
.C!IINN#J!!#

c& ',,Q/KC+%&EC%&E1vs1vWG1vs1vWG
 QU|HQT{H((+((+ 
;' 23	;' 23 9 )Z0 !/s3
	s   *?H' 'I;I

Ic                 t   ||st         t        |
|||duxr |	du      st         t        | t        j                        r| j
                  j                  nd}|dk7  sVt        j                  j                         s8t        t        j                  d      r$t        j                  j                         rt         t        |      }t        |      }t        |      }t        |      }|j                  }t        | |      } t        ||      }t        ||      }| j
                  }t        ||||      }t        |||	|      }||z   }t        j                  ||gd      }	 t        j                   j#                  | ||d	d
dt        j$                         d	d		      \  }}|j,                  d   }|dd|f   }|d|df   }|d   |d
   }}|dddf   } |dddf   }!|r|}"|}#n"| j/                  d
      }"|!j/                  d
      }#| |"|gf|!|#|gffS # t&        t(        t*        f$ r}t        d|       |d}~ww xY w)a  Fused bidirectional cuDNN GRU for the torch backend.

    Runs both directions in a single ``torch._VF.gru(bidirectional=True)``
    call. Mirrors ``bidirectional_lstm`` above. Raises ``NotImplementedError``
    on CPU, under tracing, with a mask, or with ``reset_after=False`` so the
    caller falls back to the two-pass path. Returns
    ``((fwd_last, fwd_outputs, [fwd_h_n]), (bwd_last, bwd_outputs, [bwd_h_n]))``
    with backward outputs already in original time order.
    Nr   r  r   r   r>   r   rP   Tr   r   z cuDNN bidirectional GRU failed: .r   )r   r   r   r   r   r   r   r   r   r   r   r   r   r?   r-  r   r  r  r  r9  r:  r   r   r   )$r*   fwd_initial_statebwd_initial_stater   r?  r@  rA  rB  rC  rD  r   r   r   r   r  r  r   rE  rG  r   rI  rJ  r  rK  r   r  rM  r   rN  rO  rP  rQ  rT  rU  rV  rW  s$                                       r   bidirectional_grur\  a  sw   4 {!!%>($*>	 "! )>E  	f$99!ENNN3++- "!":.J,-AB":.J,-AB$$Mv];F0FF0FF]]F#((FJ $((FJ *$F ++vv&A
.CQyy}}!!#

" ',,Q/KC+%&EC%&E1vs1vWG
 QU|HQT{H((+((+ 
;	*	;	* 3 )Z0 Q!$DQC"HIqPQs   !<H H7#H22H7)FNNFNFFT)T)FFFT)FF)FFT)r   r   r   keras.src.backend.torch.corer   r   r   r   r   r   r   r   r   r   r   r   r  r-  r  r  rX  r\  r0   r   r   <module>r^     s      : 	p,f 0F.&((52%4PM 	: Zz),X6(D Tn+4\&'R4%L !ZT vr   