o
    i e                     @   sN   d Z ddlm  mZ ddlmZ ddlmZ ddl	m
Z
 G dd de
ZdS )z0Base class for recurrent layers backed by cuDNN.    N)backend	InputSpec)RNNc                       s   e Zd ZdZ					d fdd	ZdddZ fdd	Zed
d Ze	dd Z
e	dd Ze	 fddZd fdd	Z  ZS )	_CuDNNRNNaK  Private base class for CuDNNGRU and CuDNNLSTM layers.

    Args:
      return_sequences: Boolean. Whether to return the last output
          in the output sequence, or the full sequence.
      return_state: Boolean. Whether to return the last state
          in addition to the output.
      go_backwards: Boolean (default False).
          If True, process the input sequence backwards and return the
          reversed sequence.
      stateful: Boolean (default False). If True, the last state
          for each sample at index i in a batch will be used as initial
          state for the sample of index i in the following batch.
      time_major: Boolean (default False). If true, the inputs and outputs will
          be in shape `(timesteps, batch, ...)`, whereas in the False case, it
          will be `(batch, timesteps, ...)`.
    Fc                    s   t t| jd	i | || _|| _|| _|| _|| _d| _t	ddg| _
t| jjdr0| jj}n| jjg}dd |D | _d | _d | _d| _tdg| _d S )
NF   )ndim__len__c                 S   s   g | ]	}t d |fdqS )N)shaper   ).0dim r   R/var/www/myenv/lib/python3.10/site-packages/keras/src/layers/rnn/base_cudnn_rnn.py
<listcomp>C   s    z&_CuDNNRNN.__init__.<locals>.<listcomp>r   r   )superr   __init__return_sequencesreturn_statego_backwardsstateful
time_majorsupports_maskingr   
input_spechasattrcell
state_size
state_specconstants_spec_states_num_constantstfconstant_vector_shape)selfr   r   r   r   r   kwargsr   	__class__r   r   r   ,   s    

z_CuDNNRNN.__init__Nc                 C   s   t |tr	|d }|d urtdt |tr!|dd  }|d }n|d ur&n| jr-| j}n| |}t|t| jkrPtdtt| j d tt| d | jrYt	
|d}| ||\}}| jrtdd t| j|D }| | | jr||g| S |S )	Nr   z(Masking is not supported for CuDNN RNNs.   z
Layer has z states but was passed z initial states.c                 S   s    g | ]\}}t jj||qS r   )r!   compatv1assign)r   
self_statestater   r   r   r   k   s    z"_CuDNNRNN.call.<locals>.<listcomp>)
isinstancelist
ValueErrorr   statesget_initial_statelenstrr   r   reverse_process_batchzip
add_updater   )r$   inputsmasktraininginitial_stateoutputr1   updatesr   r   r   callI   sF   







z_CuDNNRNN.callc                    sD   | j | j| j| j| jd}tt|  }tt	|
 t	|
  S )N)r   r   r   r   r   )r   r   r   r   r   r   r   
get_configdictr/   items)r$   configbase_configr&   r   r   r@   v   s   z_CuDNNRNN.get_configc                 C   s   | di |S )Nr   r   )clsrC   r   r   r   from_config   s   z_CuDNNRNN.from_configc                 C   s    | j r| jr| j| j| jgS g S N	trainablebuiltkernelrecurrent_kernelbiasr$   r   r   r   trainable_weights      z_CuDNNRNN.trainable_weightsc                 C   s    | j s| jr| j| j| jgS g S rG   rH   rN   r   r   r   non_trainable_weights   rP   z_CuDNNRNN.non_trainable_weightsc                    s   t t| jS rG   )r   r   lossesrN   r&   r   r   rR      s   z_CuDNNRNN.lossesc                    s   t t| j|dS )N)r9   )r   r   get_losses_for)r$   r9   r&   r   r   rS      s   z_CuDNNRNN.get_losses_for)FFFFF)NNNrG   )__name__
__module____qualname____doc__r   r?   r@   classmethodrF   propertyrO   rQ   rR   rS   __classcell__r   r   r&   r   r      s&    
-


r   )rW   tensorflow.compat.v2r)   v2r!   	keras.srcr   keras.src.engine.input_specr   keras.src.layers.rnn.base_rnnr   r   r   r   r   r   <module>   s   