o
    Ne                     @   sh   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 edddgdG d	d
 d
e
jZdS )z!Adagrad optimizer implementation.    N)backend_config)optimizer_v2)keras_exportzkeras.optimizers.legacy.Adagradzkeras.optimizers.Adagrad)v1c                       s|   e Zd ZdZdZ				 d fdd	Zdd	 Z fd
dZ fddZe	dddZ
dddZdddZ fddZ  ZS )Adagrada  Optimizer that implements the Adagrad algorithm.

    Adagrad is an optimizer with parameter-specific learning rates,
    which are adapted relative to how frequently a parameter gets
    updated during training. The more updates a parameter receives,
    the smaller the updates.

    Args:
      learning_rate: Initial value for the learning rate:
        either a floating point value,
        or a `tf.keras.optimizers.schedules.LearningRateSchedule` instance.
        Defaults to 0.001.
        Note that `Adagrad` tends to benefit from higher initial learning rate
        values compared to other optimizers.
        To match the exact form in the original paper, use 1.0.
      initial_accumulator_value: Floating point value.
        Starting value for the accumulators (per-parameter momentum values).
        Must be non-negative.
      epsilon: Small floating point value used to maintain numerical stability.
      name: Optional name prefix for the operations created when applying
        gradients.  Defaults to `"Adagrad"`.
      **kwargs: keyword arguments. Allowed arguments are `clipvalue`,
        `clipnorm`, `global_clipnorm`.
        If `clipvalue` (float) is set, the gradient of each weight
        is clipped to be no higher than this value.
        If `clipnorm` (float) is set, the gradient of each weight
        is individually clipped so that its norm is no higher than this value.
        If `global_clipnorm` (float) is set the gradient of all weights is
        clipped so that their global norm is no higher than this value..

    Reference:
      - [Duchi et al., 2011](
        http://www.jmlr.org/papers/volume12/duchi11a/duchi11a.pdf).
    TMbP?皙?Hz>c                    sr   |dk r
t d| |d u rt }t j|fi | | d|d| | d| j || _|p5t | _d S )Ng        z2initial_accumulator_value must be non-negative: %slearning_ratelrdecay)	
ValueErrorr   epsilonsuper__init__Z
_set_hyperget_initial_decay_initial_accumulator_value)selfr
   initial_accumulator_valuer   namekwargs	__class__ UD:\Projects\ConvertPro\env\Lib\site-packages\keras/optimizers/optimizer_v2/adagrad.pyr   E   s   zAdagrad.__init__c                 C   s8   |D ]}|j j}tjjj| j|d}| |d| qd S )Ndtypeaccumulator)r   
base_dtypetfcompatr   Zconstant_initializerr   Zadd_slot)r   Zvar_listvarr   initr   r   r   _create_slotsZ   s   zAdagrad._create_slotsc              	      sT   t  ||| |||f tt| j||||f d  tjdtjdd d S )Nlr_tr   r   )r   Zneg_lr_tzero)	r   _prepare_localupdatedictr    Zconvert_to_tensorr   ZzerosZint64)r   
var_device	var_dtypeapply_stater   r   r   r'   b   s   zAdagrad._prepare_localc                    s:   | j }t|t|d krtdg| }t | d S )N   r   )weightslennparrayr   set_weights)r   r.   paramsr   r   r   r2   l   s   zAdagrad.set_weightsNc                 C   s4   d|vrd|d< d|v r| d|d< | di |S )a  Creates an optimizer from its config.

        This method is the reverse of `get_config`,
        capable of instantiating the same optimizer from the config
        dictionary.

        Args:
            config: A Python dictionary, typically the output of get_config.
            custom_objects: A Python dictionary mapping names to additional
              Python objects used to create this optimizer, such as a function
              used for a hyperparameter.

        Returns:
            An optimizer instance.
        r   r   r   r
   Nr   )pop)clsconfigZcustom_objectsr   r   r   from_configu   s
   zAdagrad.from_configc                 C   s`   |j |jj}}|pi ||fp| ||}| |d}tjj|j	|j	|d |d || j
dS )Nr   r%   r   )r"   accumr   r   graduse_locking)devicer   r   r   _fallback_apply_stateget_slotr    raw_opsZResourceApplyAdagradV2handle_use_locking)r   r9   r"   r,   r*   r+   coefficientsaccr   r   r   _resource_apply_dense   s   
zAdagrad._resource_apply_densec           	   	   C   sb   |j |jj}}|pi ||fp| ||}| |d}tjj|j	|j	|d |d ||| j
dS )Nr   r%   r   )r"   r8   r   r   r9   indicesr:   )r;   r   r   r   r<   r=   r    r>   ZResourceSparseApplyAdagradV2r?   r@   )	r   r9   r"   rD   r,   r*   r+   rA   rB   r   r   r   _resource_apply_sparse   s    
zAdagrad._resource_apply_sparsec                    s.   t   }|| d| j| j| jd |S )Nr
   )r
   r   r   r   )r   
get_configr(   Z_serialize_hyperparameterr   r   r   )r   r6   r   r   r   rF      s   

zAdagrad.get_config)r   r   r	   r   )N)__name__
__module____qualname____doc__Z_HAS_AGGREGATE_GRADr   r$   r'   r2   classmethodr7   rC   rE   rF   __classcell__r   r   r   r   r      s     #
	

r   )rJ   numpyr0   Ztensorflow.compat.v2r!   v2r    Zkerasr   Zkeras.optimizers.optimizer_v2r   Z tensorflow.python.util.tf_exportr   ZOptimizerV2r   r   r   r   r   <module>   s   