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Zdd ZG dd	 d	ejZG d
d dejZG dd de	jZG dd dejZG dd dejZG dd dejZdS )z/A simple functional keras model with one layer.    N)model_collection_base)gradient_descent
   c                  C   sX   t jtjddt jd} t jtjddt jd}t jtjddt jd}| ||fS )Ni     dtype   )tfconstantnprandomrandfloat32)x_trainy_train	x_predict r   Q/var/www/myenv/lib/python3.10/site-packages/keras/src/distribute/simple_models.py_get_data_for_simple_models   s   
r   c                   @   (   e Zd ZdZdd Zdd Zdd ZdS )	SimpleFunctionalModelz)A simple functional model and its inputs.c                 K   s^   d}t jjdtjd}t jjdtj|d|}t j||d}tjdd}|j	d	d
g|d |S )Noutput_1)r   )shaper   r   )r   name)inputsoutputsMbP?learning_ratemsemaelossmetrics	optimizer)
keraslayersInputr	   r   DenseModelr   SGDcompile)selfkwargsoutput_namexymodelr$   r   r   r   	get_model&   s   zSimpleFunctionalModel.get_modelc                 C      t  S Nr   r,   r   r   r   get_data2      zSimpleFunctionalModel.get_datac                 C      t S r4   _BATCH_SIZEr6   r   r   r   get_batch_size5      z$SimpleFunctionalModel.get_batch_sizeN__name__
__module____qualname____doc__r2   r7   r<   r   r   r   r   r   #   s
    r   c                   @   r   )	SimpleSequentialModelz)A simple sequential model and its inputs.c                 K   sN   d}t  }t jjdtj|dd}|| tjdd}|j	ddg|d	 |S )
Nr   r   r   )r   r   	input_dimr   r   r   r    r!   )
r%   
Sequentialr&   r(   r	   r   addr   r*   r+   )r,   r-   r.   r1   r0   r$   r   r   r   r2   <   s   

zSimpleSequentialModel.get_modelc                 C   r3   r4   r5   r6   r   r   r   r7   I   r8   zSimpleSequentialModel.get_datac                 C   r9   r4   r:   r6   r   r   r   r<   L   r=   z$SimpleSequentialModel.get_batch_sizeNr>   r   r   r   r   rC   9   s
    rC   c                       s$   e Zd Z fddZdd Z  ZS )_SimpleModelc                    s"   t    tjjdtjd| _d S )Nr   r   )super__init__r%   r&   r(   r	   r   _dense_layerr6   	__class__r   r   rI   Q   s   
z_SimpleModel.__init__c                 C   s
   |  |S r4   )rJ   )r,   r   r   r   r   callU   s   
z_SimpleModel.call)r?   r@   rA   rI   rM   __classcell__r   r   rK   r   rG   P   s    rG   c                   @   r   )	SimpleSubclassModelz%A simple subclass model and its data.c                 K   s*   t  }tjdd}|jddgd|d |S )Nr   r   r   r    F)r"   r#   cloningr$   )rG   r   r*   r+   )r,   r-   r1   r$   r   r   r   r2   \   s   
zSimpleSubclassModel.get_modelc                 C   r3   r4   r5   r6   r   r   r   r7   e   r8   zSimpleSubclassModel.get_datac                 C   r9   r4   r:   r6   r   r   r   r<   h   r=   z"SimpleSubclassModel.get_batch_sizeNr>   r   r   r   r   rO   Y   s
    	rO   c                   @   s"   e Zd Zdd Zejdd ZdS )_SimpleModulec                 C   s   t d| _d S )Ng      @)r	   Variablevr6   r   r   r   rI   m   s   z_SimpleModule.__init__c                 C   s
   | j | S r4   )rS   )r,   r/   r   r   r   __call__p   s   
z_SimpleModule.__call__N)r?   r@   rA   rI   r	   functionrT   r   r   r   r   rQ   l   s    rQ   c                   @   r   )	SimpleTFModuleModelz/A simple model based on tf.Module and its data.c                 K   s
   t  }|S r4   )rQ   )r,   r-   r1   r   r   r   r2   x   s   zSimpleTFModuleModel.get_modelc                 C   r3   r4   r5   r6   r   r   r   r7   |   r8   zSimpleTFModuleModel.get_datac                 C   r9   r4   r:   r6   r   r   r   r<      r=   z"SimpleTFModuleModel.get_batch_sizeNr>   r   r   r   r   rV   u   s
    rV   )rB   numpyr   tensorflow.compat.v2compatv2r	   	keras.srcsrcr%   keras.src.distributer   keras.src.optimizers.legacyr   r;   r   ModelAndInputr   rC   r)   rG   rO   ModulerQ   rV   r   r   r   r   <module>   s   		