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

Function chain

axlearn/common/optimizers.py:81–90  ·  view source on GitHub ↗
(*args)

Source from the content-addressed store, hash-verified

79
80
81def chain(*args):
82 args = [_to_partitioned_transformation(e) for e in args]
83 base = optax.chain(*[optax.GradientTransformation(init=e.init, update=e.update) for e in args])
84
85 def partition(param_spec):
86 return tuple(e.partition(param_spec) for e in args)
87
88 return PartitionedGradientTransformation(
89 init=base.init, update=base.update, partition=partition
90 )
91
92
93def named_chain(**kwargs):

Callers 7

test_weight_scalingMethod · 0.90
sgd_optimizerFunction · 0.70
adamw_optimizerFunction · 0.70
adam_optimizerFunction · 0.70
adafactor_optimizerFunction · 0.70
lion_optimizerFunction · 0.70

Tested by 1

test_weight_scalingMethod · 0.72