o
    i e_                      @   s   d Z ddlZddlm  mZ ddlmZ ddlm	Z	 ddl
mZ ddlmZ ddlmZ ddlmZ dd	lmZ dd
lmZ edG dd de	jZdS )zAWrapper allowing a stack of RNN cells to behave as a single cell.    N)backend)
base_layer)	rnn_utils)serialization_lib)generic_utils)tf_utils)
tf_logging)keras_exportzkeras.layers.StackedRNNCellsc                       st   e Zd ZdZ fddZedd Zedd Zdd	d
ZdddZ	e
jdd Z fddZedddZ  ZS )StackedRNNCellsam  Wrapper allowing a stack of RNN cells to behave as a single cell.

    Used to implement efficient stacked RNNs.

    Args:
      cells: List of RNN cell instances.

    Examples:

    ```python
    batch_size = 3
    sentence_max_length = 5
    n_features = 2
    new_shape = (batch_size, sentence_max_length, n_features)
    x = tf.constant(np.reshape(np.arange(30), new_shape), dtype = tf.float32)

    rnn_cells = [tf.keras.layers.LSTMCell(128) for _ in range(2)]
    stacked_lstm = tf.keras.layers.StackedRNNCells(rnn_cells)
    lstm_layer = tf.keras.layers.RNN(stacked_lstm)

    result = lstm_layer(x)
    ```
    c                    sx   |D ]}dt |vrtd| dt |vrtd| q|| _|dd| _| jr1td t jdi | d S )	NcallzLAll cells must have a `call` method. Received cell without a `call` method: 
state_sizezTAll cells must have a `state_size` attribute. Received cell without a `state_size`: reverse_state_orderFzreverse_state_order=True in StackedRNNCells will soon be deprecated. Please update the code to work with the natural order of states if you rely on the RNN states, eg RNN(return_state=True). )	dir
ValueErrorcellspopr   loggingwarningsuper__init__)selfr   kwargscell	__class__r   U/var/www/myenv/lib/python3.10/site-packages/keras/src/layers/rnn/stacked_rnn_cells.pyr   <   s*   zStackedRNNCells.__init__c                 C   s0   t dd | jr| jd d d D S | jD S )Nc                 s   s    | ]}|j V  qd S N)r   ).0cr   r   r   	<genexpr>Z   s
    
z-StackedRNNCells.state_size.<locals>.<genexpr>)tupler   r   r   r   r   r   r   X   s
   zStackedRNNCells.state_sizec                 C   sP   t | jd dd d ur| jd jS t| jd jr"| jd jd S | jd jS )Nr!   output_sizer   )getattrr   r$   r   is_multiple_stater   r#   r   r   r   r$   a   s
   zStackedRNNCells.output_sizeNc              	   C   sj   g }| j r| jd d d n| jD ] }t|dd }|r%|||||d q|t|||| qt|S )Nr!   get_initial_state)inputs
batch_sizedtype)r   r   r%   appendr   #generate_zero_filled_state_for_cellr"   )r   r(   r)   r*   initial_statesr   get_initial_state_fnr   r   r   r'   j   s    z!StackedRNNCells.get_initial_statec                 K   s*  | j r| jd d d n| j}tj|tj|}g }t| j|D ]f\}	}tj|r-|n|g}t	|	dd d u}
t
|dkrD|
rD|d n|}t|	jdrR||d< n|dd  t|	r_|	jn|	j}t|	jdrw|||fd|i|\}}n|||fi |\}}|| q!|tj|tj|fS )Nr!   _is_tf_rnn_cell   r   training	constants)r   r   tfnestpack_sequence_asflattenzipr   	is_nestedr%   lenr   has_argr   r   callable__call__r+   )r   r(   statesr2   r1   r   r   nested_statesnew_nested_statesr   is_tf_rnn_cellcell_call_fnr   r   r   r      s<   
zStackedRNNCells.callc              	   C   s  t |tr	|d }dd }| jD ]n}t |tjr9|js9t|j |	| d|_W d    n1 s4w   Y  t
|dd d urE|j}nt|jrQ|jd }n|j}tj|d }tj|rrtjt|||}t|}qt|gt|  }qd| _d S )Nr   c                 S   s   t | }t| g| S r   )r3   TensorShapeas_listr"   )r)   dimshaper   r   r   get_batch_input_shape   s   z4StackedRNNCells.build.<locals>.get_batch_input_shapeTr$   )
isinstancelistr   r   Layerbuiltr   
name_scopenamebuildr%   r$   r   r&   r   r3   r4   r6   r8   map_structure	functoolspartialr"   rB   rC   )r   input_shaperF   r   
output_dimr)   r   r   r   rM      s2   





zStackedRNNCells.buildc                    sN   g }| j D ]
}|t| qd|i}t  }tt| t|  S )Nr   )	r   r+   r   serialize_keras_objectr   
get_configdictrH   items)r   r   r   configbase_configr   r   r   rT      s   

zStackedRNNCells.get_configc                 C   sB   ddl m} g }|dD ]}||||d q| |fi |S )Nr   )deserializer   )custom_objects)keras.src.layersrY   r   r+   )clsrW   rZ   deserialize_layerr   cell_configr   r   r   from_config   s   
zStackedRNNCells.from_config)NNN)NNr   )__name__
__module____qualname____doc__r   propertyr   r$   r'   r   r   shape_type_conversionrM   rT   classmethodr_   __classcell__r   r   r   r   r
   "   s    



(
 r
   )rc   rO   tensorflow.compat.v2compatv2r3   	keras.srcr   keras.src.enginer   keras.src.layers.rnnr   keras.src.savingr   keras.src.utilsr   r   tensorflow.python.platformr   r    tensorflow.python.util.tf_exportr	   rI   r
   r   r   r   r   <module>   s   