
    ij^8                         d dl 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
  ed       G d d	e             Z ed
       G d de             Zy)    N)MutableMapping)keras_export)LossScaleOptimizer)	Optimizer)serialization_libzkeras.optimizers.OptimizerMapc                   d    e Zd ZdZddZed        Zd Zd Zd Z	d Z
d	 Zd
 Zd Zedd       Zy)OptimizerMapzA class mapping variables to optimizers.

    Args:
        default_optimizer: Default Keras `Optimizer`
            for any unmapped variables.
    Nc                     t        |t              st        d|       |t        |t              st        d|       || _        |xs
 t               | _        y )Nz@default_optimizer must be a Keras Optimizer instance. Received: z.optimizer_map must be a dictionary. Received: )
isinstancer   	TypeErrordict_default_optimizer_optimizer_map)selfdefault_optimizeroptimizer_maps      y/var/www/html/emotional.easysim.app/public_html/venv/lib/python3.12/site-packages/keras/src/optimizers/multi_optimizer.py__init__zOptimizerMap.__init__   sk    +Y7./1  $Zt-L@P  #4+5tv    c                     | j                   S N)r   r   s    r   r   zOptimizerMap.default_optimizer    s    &&&r   c                 R   || j                   v r| j                   |   S | j                   j                         D cg c]  }t        j                  ||      r| }}t	        |      dkD  rt        d| d      t	        |      dk(  r| j                   |d      S | j                  S c c}w )ah  Retrieves the corresponding `Optimizer` by the string key.

        This method first attempts an exact key match. If no exact match is
        found, it treats all keys in the map as regular expression patterns
        and uses `re.fullmatch` to find a policy.

        For example,
        to apply a optimizer to all sublayers of an `encoder` block,
        the key should be explicitly set to `"encoder/.*"`. A key of
        `"encoder"` will only match the layer with that exact path.

        Args:
            key: str. The key to query for an `Optimizer`.

        Returns:
            The corresponding `Optimizer`. If no match is found, this method
            returns `self.default_optimizer`.

        Raises:
            ValueError: If the `key` matches more than one regex pattern in the
            map.

        Example:

        ```python
        >>> from keras.src import optimizers
        >>> from keras.src.optimizers.multi_optimizer import OptimizerMap
        >>> opt_adam = optimizers.Adam()
        >>> opt_sgd = optimizers.SGD()
        >>> opt_rmsprop = optimizers.RMSprop()
        >>> optimizer_map = OptimizerMap(default_optimizer=opt_rmsprop)
        >>> optimizer_map["encoder/layer_0/dense"] = opt_adam
        >>> optimizer_map["encoder/.*"] = opt_sgd
        >>> optimizer_map["decoder"] = opt_adam

        >>> optimizer_map["decoder"].name
        'adam'

        >>> optimizer_map["encoder/layer_0/dense"].name
        'adam'

        >>> optimizer_map["encoder/layer_0/attention/query"].name
        'sgd'

        >>> optimizer_map["decoder/layer_0/dense"].name
        'rmsprop'
        ```
           z)Multiple optimizers assigned to variable .r   )r   keysre	fullmatchlen
ValueErrorr   )r   keypatternmatching_keyss       r   __getitem__zOptimizerMap.__getitem__$   s    d $%%%&&s++  ..335
||GS) 
 
 }!HQOPP1$&&}Q'788%%%
s   B$c                 `   || j                   v rt        d| d| j                   |    d      t        |t              st        d|       t        |t              st        d|       t        |t
              rt        d      t        |t              rt        d      || j                   |<   y)	zSet optimizer for a given variable.

        Args:
            key: string representing the variable path or regex pattern
            optimizer: Keras Optimizer instance
        'z0' already exists in the OptimizerMap with value z.. Please make sure to not use duplicated keys.z key must be a string. Received: z8optimizer must be a Keras Optimizer instance. Received: z'optimizer cannot be LossScaleOptimizer.z#optimizer cannot be MultiOptimizer.N)r   r    r   strr   r   MultiOptimizer)r   r!   	optimizers      r   __setitem__zOptimizerMap.__setitem__e   s     $%%%C5 ,,S12 3++ 
 #s#?uEFF)Y/&K)  i!34FGGi0BCC#,C r   c                     | j                   |= y r   )r   )r   r!   s     r   __delitem__zOptimizerMap.__delitem__   s    $r   c                 ,    t        | j                        S r   )r   r   r   s    r   __len__zOptimizerMap.__len__   s    4&&''r   c                 ,    t        | j                        S r   )iterr   r   s    r   __iter__zOptimizerMap.__iter__   s    D''((r   c                      | |j                      S r   )path)r   variables     r   __call__zOptimizerMap.__call__   s    HMM""r   c                     t        j                  | j                        t        j                  | j                        dS )N)r   r   )r   serialize_keras_objectr   r   r   s    r   
get_configzOptimizerMap.get_config   s<    !2!I!I''" /EE##	
 	
r   c                     t        j                  |d   |      }t        j                  |d   |      } | ||      }|S )Nr   custom_objectsr   )r   deserialize_keras_object)clsconfigr;   r   r   obj_maps         r   from_configzOptimizerMap.from_config   sN    -FF&'
 *BB?#N
 '7r   r   )__name__
__module____qualname____doc__r   propertyr   r$   r*   r,   r.   r1   r5   r8   classmethodr@    r   r   r	   r	   
   sX    6 ' '?&B-4%()#
  r   r	   zkeras.optimizers.MultiOptimizerc                        e Zd ZdZd fd	Z fdZed        Zed        Zed        Z	e	j                  d        Z	ed        Zej                  d	        Zed
        ZddZd Zd Zd Zedd       Z xZS )r(   a  An optimizer wrapper that delegates variables to different optimizers.

    Initialize the object with an OptimizerMap instance or a callable
    function that returns an optimizer for a given variable.

    Example:
        model.compile(
            optimizer=MultiOptimizer(
                OptimizerMap(default_optimizer=optimizers.SGD(),
                {"encoder/.*": optimizers.Adam()})
            ),
            loss="binary_crossentropy",
        )

        # Or using a callable
        def optimizer_selector(variable):
            if "encoder" in variable.path:
                return optimizers.Adam()
            else:
                return optimizers.SGD()

        model.compile(
            optimizer=MultiOptimizer(optimizer_selector),
            loss="binary_crossentropy",
        )

    To access the attributes of the sub-optimizers, iterate over the
    optimizers using `.optimizers`:

    For example:

    optimizer = MultiOptimizer(OptimizerMap(
        default_optimizer=optimizers.Adam()
    ))
    optimizer['.encoder'] = optimizers.SGD()

    for optim in optimizer.optimizers:
        print(optim.learning_rate)
        print(optim.iterations)
        print(optim.loss_scale_factor)
        ...

    The MultiOptimizer class instances will not expose `learning_rate`
    attribute and will raise an error if accessed. This is because the
    learning rate might be different for different sub-optimizers.

    Note: Optimizer-specific callbacks are not supported yet.

    c                    t         |   d||       || _        g | _        t	        | j                  d      rP| j                  j                         D ]3  }|| j                  vs||_        | j                  j                  |       5 t        | j                  dd      }|2|| j                  vr#||_        | j                  j                  |       yyy)ac  
        Initialize the MultiOptimizer.

        Args:
            optimizer_map: An OptimizerMap instance or a callable function that
                returns an optimizer for a given variable.
            loss_scale_factor: It overrides the loss_scale_factor passed
            to the sub-optimizers.
            name: The name of the optimizer.
        g        )learning_rateloss_scale_factornamevaluesr   N)	superr   r   _inner_optimizershasattrrM   rK   appendgetattr)r   r   rK   rL   optdefault_opt	__class__s         r   r   zMultiOptimizer.__init__   s     	/ 	 	

 ,!#4&&1**113 7d444,=C)**11#67
 d113FM#4#9#99,=K)""))+6 : $r   c                    | j                   ry i | _        i | _        t        t	        | j
                              D cg c]  }g  }}t        |      D ]"  \  }}| j                  |      }t        |t              st        d| d      t        |t              rt        d      t        |t              rt        d      || j
                  vr=| j                  |_        | j
                  j                  |       |j                  g        | j
                  j                  |      }|| j                  | j!                  |      <   || j                  | j!                  |      <   ||   j                  |       % t#        | j
                  |      D ]  \  }}|s	|j%                  |        t&        	| I  |       y c c}w )NzOptimizer for variable z is not an Optimizer instance.z;LossScaleOptimizer cannot be used inside an MultiOptimizer.z7MultiOptimizer cannot be used inside an MultiOptimizer.)built_var_to_optimizer_idx_trainable_variables_indicesranger   rO   	enumerater   r   r   r    r   r(   rK   rQ   index_var_keyzipbuildrN   )
r   var_list_optimizer_varsivarrS   idx	variablesrU   s
            r   r_   zMultiOptimizer.build   s   ::%'",.)&+C0F0F,G&HI"II) 	,FAs%%c*Cc9- -cU 3- -  #12 &  #~. M  $000(,(>(>%&&--c2%%b)((..s3C=@D&&t}}S'9:DED--dmmC.@A3&&s+1	,6 "$"8"8.I 	%NC		)$	% 	hC Js   	F;c                     | j                   S r   )rO   r   s    r   
optimizerszMultiOptimizer.optimizers  s    %%%r   c                 |    | j                   d d  }| j                  D ]  }|j                  |j                          |S r   )
_variablesrO   extendrf   )r   varsrS   s      r   rf   zMultiOptimizer.variables#  s<     q!)) 	'CKK&	'r   c                     t        d      )NzxLearning rate cannot be accessed on a MultiOptimizer. Access the learning rate on the individual sub-optimizers instead.AttributeErrorr   s    r   rJ   zMultiOptimizer.learning_rate+  s    Q
 	
r   c                     t        d      )NzpLearning rate cannot be set on a MultiOptimizer. Set the learning rate on the individual sub-optimizers instead.rn   )r   values     r   rJ   zMultiOptimizer.learning_rate2  s    N
 	
r   c                     t        | dd       S )N_loss_scale_factor)rR   r   s    r   rK   z MultiOptimizer.loss_scale_factor9  s    t1488r   c                 \    || _         t        | d      r| j                  D ]	  }||_         y y )NrO   )rs   rP   rO   rK   )r   rq   rS   s      r   rK   z MultiOptimizer.loss_scale_factor=  s7    "'4,--- .(-%. .r   c                     | j                   S r   )_iterationsr   s    r   
iterationszMultiOptimizer.iterationsD  s    r   c                    t        |      dk(  ry |#| j                  st        d      | j                  }| j                  s| j	                  |       d| _        t        |      t        | j
                        k7  r.t        dt        |       dt        | j
                         d      t        t        | j                              D cg c]  }g  }}t        ||      D ]9  \  }}| j                  | j                  |         }||   j                  ||f       ; t        | j                  |      D ]7  \  }}	|	s	t        |	 \  }
}|j                  t        |
      t        |             9 | j                  j                  d       y c c}w )Nr   zWhen passing `grads` without `variables`, the optimizer must already be built on a list of variables. Call `optimizer.build(trainable_variables)` first.Tz>Gradients must match trainable variables one-to-one. Received z gradients and z variables.r   )r   rW   r    _trainable_variablesr_   rY   rZ   rO   r^   rX   r]   rQ   applylistrv   
assign_add)r   gradstrainable_variablesra   grads_and_varsgradrd   re   rS   sub_grads_and_vars	sub_gradssub_varss               r   rz   zMultiOptimizer.applyH  s   u:?&:: I 
 #'";";zzJJ*+DJu:T>>??J<t889:+G  ',C0F0F,G&HI"IIU$78 	4ID#,,T]]3-?@C3&&c{3	4
 (+""N(
 	;#C# "&)+=&>#	8		$y/4>:	; 	##A& Js   	Fc                 :   t        t        | j                              D cg c]  }g  }}|D ]4  }| j                  | j	                  |         }||   j                  |       6 t        | j                  |      D ]  \  }}|s	|j                  |        y c c}w r   )rZ   r   rO   rX   r]   rQ   r^   finalize_variable_values)r   r`   ra   rb   rd   re   rS   rf   s           r   r   z'MultiOptimizer.finalize_variable_valueso  s    &+C0F0F,G&HI"II 	,C,,T]]3-?@C3&&s+	, "$"8"8.I 	8NC,,Y7	8 Js   	Bc                 V   | j                   st        d      t        | j                        }t	        |      D ]#  }| j                  |   j                  ||          % |}| j                  D ];  }t        |j                        }|dkD  s||||z    }|j                  |       ||z  }= y )NzYou are calling `set_weights()` on an optimizer that has not yet been built. Please call `optimizer.build(trainable_variables)` to create the optimizer weights before calling `set_weights()`.r   )	rW   r    r   rj   rZ   assignrO   rf   set_weights)r   weightsown_var_countrc   re   rS   num_opt_varsopt_weightss           r   r   zMultiOptimizer.set_weightsy  s    zzD  DOO,}% 	2AOOA%%gaj1	2 )) 	$Cs}}-La%cC,,>?,|#	$r   c                 p    t        j                  | j                        | j                  | j                  dS )N)r   rK   rL   )r   r7   r   rK   rL   r   s    r   r8   zMultiOptimizer.get_config  s5    .EE## "&!7!7II
 	
r   c                 l    |j                         }t        j                  |d   |      |d<    | di |S )Nr   r:   rG   )copyr   r<   )r=   r>   r;   s      r   r@   zMultiOptimizer.from_config  s:    "3"L"L?#N#
 }V}r   )NNr   )rA   rB   rC   rD   r   r_   rE   rh   rf   rJ   setterrK   rw   rz   r   r   r8   rF   r@   __classcell__)rU   s   @r   r(   r(      s    0d7B' R & &   
 
 
 
 9 9 . .    %'N8$*
  r   r(   )r   collections.abcr   keras.src.api_exportr   )keras.src.optimizers.loss_scale_optimizerr   keras.src.optimizers.optimizerr   keras.src.savingr   r	   r(   rG   r   r   <module>r      s`    	 * - H 4 . -.S> S /Sl /0{Y { 1{r   