MCPcopy Create free account
hub / github.com/pytorch/pytorch / fused_adam

Method fused_adam

test/test_proxy_tensor.py:810–830  ·  view source on GitHub ↗
(params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs, state_steps)

Source from the content-addressed store, hash-verified

808 state_steps = [torch.tensor(0) for _ in range(10)]
809
810 def fused_adam(params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs, state_steps):
811 (new_params, _, _, _, _) = aten._fused_adam.default(
812 params,
813 grads,
814 exp_avgs,
815 exp_avg_sqs,
816 max_exp_avg_sqs,
817 state_steps,
818 lr=0.1,
819 beta1=0.9,
820 beta2=0.999,
821 weight_decay=0.01,
822 eps=1e-8,
823 amsgrad=False,
824 maximize=False,
825 )
826
827 for p, new_p in zip(params, new_params):
828 p.copy_(new_p)
829
830 return params
831
832 gm = make_fx(fused_adam, tracing_mode='fake')(
833 params,

Callers

nothing calls this directly

Calls 2

defaultMethod · 0.80
copy_Method · 0.45

Tested by

no test coverage detected