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	l	mZ dd
l	mZ ddl	mZ G dd dejZdS )zNadam optimizer implementation.    )ops)tensor_conversion)backend_config)learning_rate_schedule)optimizer_v2)	array_ops)control_flow_ops)math_ops)	state_ops)	variablesc                       sl   e Zd ZdZdZ					 d fdd	Zd	d
 Zdd Z fddZdddZ	dddZ
 fddZ  ZS )Nadama  Optimizer that implements the NAdam algorithm.
  Much like Adam is essentially RMSprop with momentum, Nadam is Adam with
  Nesterov momentum.

  Args:
    learning_rate: A Tensor or a floating point value.  The learning rate.
    beta_1: A float value or a constant float tensor. The exponential decay
      rate for the 1st moment estimates.
    beta_2: A float value or a constant float tensor. The exponential decay
      rate for the exponentially weighted infinity norm.
    epsilon: A small constant for numerical stability.
    name: Optional name for the operations created when applying gradients.
      Defaults to `"Nadam"`.
    **kwargs: Keyword arguments. Allowed to be one of
      `"clipnorm"` or `"clipvalue"`.
      `"clipnorm"` (float) clips gradients by norm; `"clipvalue"` (float) clips
      gradients by value.

  Usage Example:
    >>> opt = tf.keras.optimizers.Nadam(learning_rate=0.2)
    >>> var1 = tf.Variable(10.0)
    >>> loss = lambda: (var1 ** 2) / 2.0
    >>> step_count = opt.minimize(loss, [var1]).numpy()
    >>> "{:.1f}".format(var1.numpy())
    9.8

  Reference:
    - [Dozat, 2015](http://cs229.stanford.edu/proj2015/054_report.pdf).
  TMbP??+?Hz>c                    s   | dd|d< |d|}t|tjrtdtt| j|fi | | 	d|d| | 	d| j
 | 	d| | 	d| |pFt | _d | _d S )	Nschedule_decaygMbp?decaylrzdThe Nadam optimizer does not support tf.keras.optimizers.LearningRateSchedules as the learning rate.learning_ratebeta_1beta_2)popget
isinstancer   LearningRateSchedule
ValueErrorsuperr   __init__
_set_hyper_initial_decayr   epsilon_m_cache)selfr   r   r   r    namekwargs	__class__ Y/var/www/myenv/lib/python3.10/site-packages/tensorflow/python/keras/optimizer_v2/nadam.pyr   ?   s   
zNadam.__init__c                 C   sp   |d j j}| jd u r | jdg |ddtjjd| _| j| j |D ]}| 	|d q"|D ]}| 	|d q-d S )Nr   momentum_cacheonesF)shapedtypeinitializer	trainableaggregationmv)
r,   
base_dtyper!   
add_weighttf_variablesVariableAggregationONLY_FIRST_REPLICA_weightsappendadd_slot)r"   var_list	var_dtypevarr'   r'   r(   _create_slotsV   s    
zNadam._create_slotsc                 C   s<  t | d|}t | d|}t | d|}t| jd |}t| jd |}td|}	|ddt|	| j|    }
|ddt|	| j|    }t| j||
 }|| j	j
u rmt tj| j	|| jd	}|| }t|| t| j||||
|d| d| d|
 d| d| dt|| d
|||f< d S )Nr   r   r         gQ?g      ?g      ?use_locking)lr_tneg_lr_tr    beta_1_tbeta_2_tm_tm_t_1one_minus_beta_1_tone_minus_beta_2_tone_minus_m_tone_minus_m_schedule_newone_minus_m_schedule_nextv_t_prime_denominator)r   identity
_get_hyperr	   cast
iterationspowr   _m_cache_readr!   r,   r
   assign_use_lockingdictr   "convert_to_tensor_v2_with_dispatchr    )r"   
var_devicer;   apply_staterB   rD   rE   
local_step	next_step
decay_baserF   rG   m_schedule_newm_schedule_nextr'   r'   r(   _prepare_locali   sF   
zNadam._prepare_localc                    s   t | j| _tt| |S N)r   rN   r!   rS   r   r   _prepare)r"   r:   r%   r'   r(   ra      s   zNadam._prepareNc                 C   s  |j |jj}}|pi ||fp| ||}| |d}| |d}||d  }	|d | |d |  }
tj||
| jd}
|
|d  }|d | |d	 t	
|  }tj||| jd}||d
  }|d |	 |d |  }||d | t	||d    }tj||| jdjS )Nr0   r1   rK   rD   rH   r@   rL   rE   rI   rM   rJ   rG   rB   r    )devicer,   r2   r   _fallback_apply_stateget_slotr
   rT   rU   r	   squaresqrtop)r"   gradr<   rY   rX   r;   coefficientsr0   r1   g_primerF   	m_t_primev_t	v_t_primem_t_barvar_tr'   r'   r(   _resource_apply_dense   s0   





zNadam._resource_apply_densec                 C   s  |j |jj}}|pi ||fp| ||}| |d}| |d}	||d  }
||d  }tj|||d  | jd}t	
|g | |||}t||}W d    n1 sZw   Y  ||d  }|d |
 |d	 |  }|| |d
  }tj|	|	|d  | jd}t	
|g | |	||}t||}W d    n1 sw   Y  ||d  }t||d  }| |||d | | }tj|||g S )Nr0   r1   rK   rH   rD   r@   rL   rJ   rG   rI   rE   rM   r    rC   )rb   r,   r2   r   rc   rd   r
   rT   rU   r   control_dependencies_resource_scatter_addr   gatherr	   rf   r   group)r"   rh   r<   indicesrY   rX   r;   ri   r0   r1   rj   m_scaled_g_valuesrF   	m_t_slicerk   rn   v_scaled_g_valuesrl   	v_t_slicerm   v_prime_sqrt_plus_eps
var_updater'   r'   r(   _resource_apply_sparse   sD   


zNadam._resource_apply_sparsec                    s>   t t|  }|| d| j| d| d| jd |S )Nr   r   r   )r   r   r   r   r    )r   r   
get_configupdate_serialize_hyperparameterr   r    )r"   configr%   r'   r(   r}      s   zNadam.get_config)r   r   r   r   r   r`   )__name__
__module____qualname____doc___HAS_AGGREGATE_GRADr   r=   r_   ra   rp   r|   r}   __classcell__r'   r'   r%   r(   r      s    &

(r   N)r   tensorflow.python.frameworkr   r   tensorflow.python.kerasr   $tensorflow.python.keras.optimizer_v2r   r   tensorflow.python.opsr   r   r	   r
   r   r4   OptimizerV2r   r'   r'   r'   r(   <module>   s   