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 dZdd Zdd Zdd Zdd Zdd Zejdd ZdS )zE2E test for DTensor with Mnist model.

Note that this is used as prototype and verification of current functionality,
and will be changed rapidly. Please don't reply on any of these methods as a
public API/contract.
    N)logging)layers)losses)models)mnist)dtensor_api
layout_map)np_utils
   c                 C   s4   t |  t W  d   S 1 sw   Y  dS )zBuilds a Sequential CNN model to recognize MNIST digits.

    Args:
      layout_map: dict of string name -> Layout, for weights creation.

    Returns:
      a CNN Keras model used for MNIST
    N)layout_map_liblayout_map_scope	get_modelr    r   W/var/www/myenv/lib/python3.10/site-packages/keras/src/dtensor/integration_test_utils.pyget_model_with_layout_map&   s   
$r   c               	   C   s   t  } | tjdddddd | tjddddd	 | tjd
d | td | t  | tjdddd | td | tjt	ddd | S )z8Builds a Sequential CNN model to recognize MNIST digits.    conv2d_1)   r   relu)   r      )namekernel_size
activationinput_shape@   conv2d_2)r   r   r   )   r   )	pool_sizeg      ?   dense_1)r   r   g      ?dense_2softmax)
r   
Sequentialaddr   Conv2DMaxPooling2DDropoutFlattenDense	NUM_CLASS)modelr   r   r   r   5   sJ   	r   c                 C   s`   t j| d}tjj| dd}tjj| dd}tjj| dd}||d< ||d< ||d< ||d	< |S )
N)mesh   rankr   r   zconv2d.*kernelzconv2d.*biaszdense.*kernelzdense.*bias)r   	LayoutMapdtensorLayout
replicated)r-   r	   	layout_4d	layout_2d	layout_1dr   r   r   get_all_replicated_layout_map^   s   r8   c                 C   s   t  \\}}\}}tj|ddd}tj|ddd}|d }|d }t|| }t|| }tjj	
||f j|dd}tjj	
||f j|dd}||fS )N)axisfloat32   T)drop_remainder)r   	load_datanpexpand_dimsastyper
   to_categoricaltfdataDatasetfrom_tensor_slicesrepeatbatch)	num_class
batch_sizex_trainy_trainx_testy_testtrain_dseval_dsr   r   r   get_mnist_datasetsm   s$   rQ   c              	   C   s   t t|\}}tjj|ddd}tjj|ddd}	t }
| }t|}g }t	|D ]F}d}t	|D ]*}t
|\}}t||}t||}t||}t||	}|t| |||
|7 }q3t|| }td|| || q+|S )NrH   r.   r/   r   g        zEpoch %d, Loss: %f)rQ   r+   r2   r3   batch_shardedr   CategoricalCrossentropynum_local_devicesiterrangenextrC   splitpack
train_stepreduce_meanr   infoappend)r,   	optimizerr-   
num_epochssteps_per_epochglobal_batch_sizedataset_input_image_layoutinput_label_layoutloss_objrT   iteratortrain_lossesepoch
total_lossimageslabelsd_imagesd_labels
train_lossr   r   r   train_mnist_model_batch_sharded   s,   
rp   c           	      C   sb   t  }| |dd}|||}W d    n1 sw   Y  ||| j}|t|| j |S )NT)training)rC   GradientTapegradienttrainable_variablesapply_gradientszip)	r,   featurelabelrf   r^   tapepredictloss	gradientsr   r   r   rZ      s   
rZ   )__doc__numpyr?   tensorflow.compat.v2compatv2rC   abslr   	keras.srcr   r   r   keras.src.datasetsr   keras.src.dtensorr   r2   r	   r   keras.src.utilsr
   r+   r   r   r8   rQ   rp   functionrZ   r   r   r   r   <module>   s&   )"