| 91 | |
| 92 | |
| 93 | def named_chain(**kwargs): |
| 94 | kwargs = {k: _to_partitioned_transformation(v) for k, v in kwargs.items()} |
| 95 | |
| 96 | def init_fn(params): |
| 97 | return {k: v.init(params) for k, v in kwargs.items()} |
| 98 | |
| 99 | def update_fn( |
| 100 | updates: NestedTensor, state: dict[str, Any], params: NestedOptParam |
| 101 | ) -> tuple[NestedTensor, optax.EmptyState]: |
| 102 | new_state = {} |
| 103 | for k, v in kwargs.items(): |
| 104 | updates, new_state[k] = v.update(updates, state[k], params) |
| 105 | return updates, new_state |
| 106 | |
| 107 | def partition_fn(param_spec): |
| 108 | return {k: v.partition(param_spec) for k, v in kwargs.items()} |
| 109 | |
| 110 | return PartitionedGradientTransformation(init=init_fn, update=update_fn, partition=partition_fn) |
| 111 | |
| 112 | |
| 113 | def _no_op(): |