(self)
| 890 | def test_cast_floats_per_param(self, per_param_train_dtype, in_tree, casted_tree): |
| 891 | out_tree = cast_floats_per_param(in_tree, per_param_train_dtype) |
| 892 | jax.tree.map(self.assertTensorEqual, out_tree, casted_tree) |
| 893 | |
| 894 | def test_count_model_params(self): |
| 895 | tree = { |
| 896 | "a": jnp.asarray([1]), |
| 897 | "b": [jnp.asarray([2]), {"c": jnp.asarray([3]), "d": jnp.asarray([4])}], |
| 898 | "e": None, |
| 899 | } |
| 900 | self.assertEqual(4, count_model_params(tree)) |
| 901 |
nothing calls this directly
no test coverage detected