Method
fused_adam
(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
Tested by
no test coverage detected