o
    i e                     @   s|   d 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
dZG dd dejZd	S )z0The implementation of `tf.data.Dataset.shuffle`.    )dataset_ops)structured_function)nest)	structure)ops)gen_experimental_dataset_ops)collections_abcNc                 C   s   t | ||||dS )N)name)_ScanDataset)input_datasetinitial_state	scan_funcuse_default_devicer	    r   Q/var/www/myenv/lib/python3.10/site-packages/tensorflow/python/data/ops/scan_op.py_scan   s   
r   c                       sB   e Zd ZdZ		d fdd	Zdd Zedd Zd	d
 Z  Z	S )r
   z1A dataset that scans a function across its input.Nc                    s  || _ t|| _t| j| _d}|rtj||  | j|j	fdd}t
|jtjr1t|jdks:td|j d|j\}| _|j\}}	tdd | j}
tt|t|
D ]\}}t||smtd	|
 d
| dqY|j\}}tdd | j}tt|t|D ]\}}||krtd| d
| dq|j\}}tdd | j}t|||	| _t|}t|}dd t||D }d}t||D ]\}}|jdur|jdu s| | krd} nq|rt|t|||
| _|s|| _| jj t!"  || _#|dur)t$j%| j j&t'| j| j| jjj(f| jjd|d| j)}nt$j%| j j&t'| j| j| jjj(f| jjdd| j)}t* +|| dS )zSee `scan()` for details.TF)input_structureadd_to_graph   zzInvalid `scan_func`. `scan_func` should return a pair consisting of new state and the output value but its return type is .c                 S      |   S N)_to_legacy_output_classescomponent_specr   r   r   <lambda>J       z'_ScanDataset.__init__.<locals>.<lambda>zbInvalid `scan_func`. The element classes for the new state must match the initial state. Expected z, got c                 S   r   r   )_to_legacy_output_typesr   r   r   r   r   V   r   z`Invalid `scan_func`. The element types for the new state must match the initial state. Expected c                 S   r   r   )_to_legacy_output_shapesr   r   r   r   r   b   r   c                 S   s   g | ]	\}}| |qS r   )most_specific_compatible_shape).0originalnewr   r   r   
<listcomp>i   s    z)_ScanDataset.__init__.<locals>.<listcomp>N)fpreserve_cardinalityr   )r$   r%   ),_input_datasetr   normalize_element_initial_statetype_spec_from_value_state_structurer   StructuredFunctionWrapper_transformation_nameelement_spec
isinstanceoutput_typesr   Sequencelen	TypeErroroutput_structureoutput_classes_output_classesr   map_structurezipflatten
issubclassoutput_shapesconvert_legacy_structure_element_specndimsas_listpack_sequence_as
_scan_funcfunctionr   r   get_default_graph_nameged_opsscan_dataset_variant_tensorto_tensor_listcaptured_inputs_common_argssuper__init__)selfr   r   r   r   r	   need_to_rerunwrapped_funcnew_state_classesr4   old_state_classesnew_state_classold_state_classnew_state_typesr/   old_state_typesnew_state_typeold_state_typenew_state_shapesr:   old_state_shapesflat_state_shapesflat_new_state_shapesweakened_state_shapesoriginal_shapeweakened_shapevariant_tensor	__class__r   r   rK   %   s   












I
	z_ScanDataset.__init__c                 C   s   | j gS r   )r@   rL   r   r   r   
_functions   s   z_ScanDataset._functionsc                 C   s   | j S r   )r<   ra   r   r   r   r-      s   z_ScanDataset.element_specc                 C   s   dS )NzDataset.scan()r   ra   r   r   r   r,      s   z!_ScanDataset._transformation_nameNN)
__name__
__module____qualname____doc__rK   rb   propertyr-   r,   __classcell__r   r   r_   r   r
   "   s    s
r
   rc   )rg   tensorflow.python.data.opsr   r   tensorflow.python.data.utilr   r   tensorflow.python.frameworkr   tensorflow.python.opsr   rD   tensorflow.python.util.compatr   r   UnaryDatasetr
   r   r   r   r   <module>   s   
	