o
    i e1                     @   sZ   d Z ddlm  mZ ddlmZ ddlmZ ddlm	Z	 ddd	Z
	
			dddZdS )z.A simple network to use in tests and examples.    N)core)normalization)optimizer_v2FTc                    s0   dd }t jd|d  fdd}|| fS )z.Example of non-distribution-aware legacy code.c                  S   s$   t jjdgg } | jdddS )N      ?   T)drop_remainder)tfdataDatasetfrom_tensorsrepeatbatch)dataset r   P/var/www/myenv/lib/python3.10/site-packages/keras/src/distribute/test_example.py
dataset_fn   s   z)minimize_loss_example.<locals>.dataset_fnr   use_biasc                    sH    fdd}t tjr|fddS r|S | S )z(A very simple model written by the user.c                     s"   t  g t d } | |  S )Nr   )r   reshapeconstant)y)layerxr   r   loss_fn&   s   z8minimize_loss_example.<locals>.model_fn.<locals>.loss_fnc                          j S Ntrainable_variablesr   r   r   r   <lambda>,       z9minimize_loss_example.<locals>.model_fn.<locals>.<lambda>
isinstancer   OptimizerV2minimizer   r   r   	optimizeruse_callable_lossr   r   model_fn#   s   
z'minimize_loss_example.<locals>.model_fn)r   Dense)r'   r   r(   r   r*   r   r&   r   minimize_loss_example   s   
r,   r   ?c                    sL    fdd}|  t j||ddtjdddfdd}||fS )	zKExample of non-distribution-aware legacy code with batch
    normalization.c                      s    t jjdd t D  S )Nc                    s"   g | ]  fd dt dD qS )c                    s$   g | ]  fd dt dD qS )c                    s$   g | ]}t  d  | d  qS )   d   )float).0r   )r   zr   r   
<listcomp>F   s   $ zObatchnorm_example.<locals>.dataset_fn.<locals>.<listcomp>.<listcomp>.<listcomp>r.   ranger1   r2   r)   r   r3   E   s    zDbatchnorm_example.<locals>.dataset_fn.<locals>.<listcomp>.<listcomp>   r4   r6   r   r7   r   r3   D   s    
z9batchnorm_example.<locals>.dataset_fn.<locals>.<listcomp>)r   r	   r
   from_tensor_slicesr5   r   r   )batch_per_epochr   r   r   @   s   z%batchnorm_example.<locals>.dataset_fnF)renormmomentumfusedr   r   c                    s<    fdd}t tjr|fddS |S )zA model that uses batchnorm.c                     st    dd} t rt jjt jjjjng  t t | t 	d }W d    |S 1 s3w   Y  |S )NT)trainingr   )
r   control_dependenciescompatv1get_collection	GraphKeys
UPDATE_OPSreduce_mean
reduce_sumr   )r   loss)	batchnormr   update_ops_in_replica_moder   r   r   r   V   s   


z4batchnorm_example.<locals>.model_fn.<locals>.loss_fnc                      r   r   r   r   r   r   r   r   f   r    z5batchnorm_example.<locals>.model_fn.<locals>.<lambda>r!   r%   )rH   r   r'   rI   r)   r   r*   S   s   
z#batchnorm_example.<locals>.model_fn)r   BatchNormalizationr   r+   )optimizer_fnr:   r<   r;   rI   r   r*   r   )r:   rH   r   r'   rI   r   batchnorm_example6   s   

rL   )FT)r   r-   FF)__doc__tensorflow.compat.v2r@   v2r   keras.src.legacy_tf_layersr   r   keras.src.optimizers.legacyr   r,   rL   r   r   r   r   <module>   s   
 