@@ -244,6 +244,12 @@ def test_map_non_params_to_none(self):
244244 def test_tree_get_all_with_path (self ):
245245 params = jnp .array ([1.0 , 2.0 , 3.0 ])
246246
247+ with self .subTest ('Test with flat tree' ):
248+ tree = ()
249+ self .assertRaises (ValueError , _state_utils .tree_get , tree , 'foo' )
250+ tree = jnp .array ([1.0 , 2.0 , 3.0 ])
251+ self .assertRaises (ValueError , _state_utils .tree_get , tree , 'foo' )
252+
247253 with self .subTest ('Test with single value in state' ):
248254 key = 'count'
249255 opt = transform .scale_by_adam ()
@@ -253,8 +259,8 @@ def test_tree_get_all_with_path(self):
253259 self .assertEqual (values_found , expected_result )
254260
255261 with self .subTest ('Test with no value in state' ):
256- key = 'count '
257- opt = alias .sgd (learning_rate = 1.0 )
262+ key = 'apple '
263+ opt = alias .adam (learning_rate = 1.0 )
258264 state = opt .init (params )
259265 values_found = _state_utils .tree_get_all_with_path (state , key )
260266 self .assertEmpty (values_found )
@@ -318,7 +324,7 @@ def test_tree_get(self):
318324
319325 with self .subTest ('Test jitted tree_get' ):
320326 opt = _inject .inject_hyperparams (alias .sgd )(
321- learning_rate = lambda x : 1 / ( x + 1 )
327+ learning_rate = lambda x : 1 / ( x + 1 )
322328 )
323329 state = opt .init (params )
324330
@@ -327,10 +333,64 @@ def get_learning_rate(state):
327333 return _state_utils .tree_get (state , 'learning_rate' )
328334
329335 for i in range (4 ):
330- # we simply update state, we don't care about updates.
336+ # we simply update state, we don't care about updates.
331337 _ , state = opt .update (params , state )
332338 lr = get_learning_rate (state )
333- self .assertEqual (lr , 1 / (i + 1 ))
339+ self .assertEqual (lr , 1 / (i + 1 ))
340+
341+ def test_tree_set (self ):
342+ params = jnp .array ([1.0 , 2.0 , 3.0 ])
343+
344+ with self .subTest ('Test with flat tree' ):
345+ tree = ()
346+ self .assertRaises (ValueError , _state_utils .tree_get , tree , 'foo' )
347+ tree = jnp .array ([1.0 , 2.0 , 3.0 ])
348+ self .assertRaises (ValueError , _state_utils .tree_get , tree , 'foo' )
349+
350+ with self .subTest ('Test modifying an injected hyperparam' ):
351+ opt = _inject .inject_hyperparams (alias .adam )(learning_rate = 1.0 )
352+ state = opt .init (params )
353+ new_state = _state_utils .tree_set (state , learning_rate = 2.0 , b1 = 3.0 )
354+ lr = _state_utils .tree_get (new_state , 'learning_rate' )
355+ self .assertEqual (lr , 2.0 )
356+
357+ with self .subTest ('Test modifying an attribute of the state' ):
358+ opt = _inject .inject_hyperparams (alias .adam )(learning_rate = 1.0 )
359+ state = opt .init (params )
360+ new_state = _state_utils .tree_set (state , learning_rate = 2.0 , b1 = 3.0 )
361+ b1 = _state_utils .tree_get (new_state , 'b1' )
362+ self .assertEqual (b1 , 3.0 )
363+
364+ with self .subTest ('Test modifying a value not present in the state' ):
365+ opt = _inject .inject_hyperparams (alias .adam )(learning_rate = 1.0 )
366+ state = opt .init (params )
367+ self .assertRaises (KeyError , _state_utils .tree_set , state , ema = 2.0 )
368+
369+ with self .subTest ('Test jitted tree_set' ):
370+
371+ @jax .jit
372+ def set_learning_rate (state , lr ):
373+ return _state_utils .tree_set (state , learning_rate = lr )
374+
375+ modified_state = state
376+ lr = 1.0
377+ for i in range (4 ):
378+ modified_state = set_learning_rate (modified_state , lr / (i + 1 ))
379+ # we simply update state, we don't care about updates.
380+ _ , modified_state = opt .update (params , modified_state )
381+ modified_lr = _state_utils .tree_get (modified_state , 'learning_rate' )
382+ self .assertEqual (modified_lr , lr / (i + 1 ))
383+
384+ with self .subTest ('Test modifying several values at once' ):
385+ opt = combine .chain (
386+ alias .adam (learning_rate = 1.0 ), alias .adam (learning_rate = 1.0 )
387+ )
388+ state = opt .init (params )
389+ new_state = _state_utils .tree_set (state , count = 2.0 )
390+ values_found = _state_utils .tree_get_all_with_path (new_state , 'count' )
391+ self .assertLen (values_found , 2 )
392+ for _ , value in values_found :
393+ self .assertEqual (value , 2.0 )
334394
335395
336396def _fake_params ():
0 commit comments