MCPcopy Create free account
hub / github.com/apple/axlearn / named_chain

Function named_chain

axlearn/common/optimizers.py:93–110  ·  view source on GitHub ↗
(**kwargs)

Source from the content-addressed store, hash-verified

91
92
93def 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
113def _no_op():

Calls 3

itemsMethod · 0.80