| 206 | } _invoker_initializer; |
| 207 | |
| 208 | void invoke(std::shared_ptr<Adam_Optimizer> opt, |
| 209 | torch::Tensor& params, |
| 210 | torch::Tensor& grads, |
| 211 | torch::Tensor& exp_avg, |
| 212 | torch::Tensor& exp_avg_sq, |
| 213 | size_t param_size, |
| 214 | bool parallel = true) |
| 215 | { |
| 216 | c10::ScalarType params_type = at::typeMetaToScalarType(params.options().dtype()); |
| 217 | c10::ScalarType state_type = at::typeMetaToScalarType(exp_avg.options().dtype()); |
| 218 | |
| 219 | auto it = invokers.find(std::tuple(params_type, state_type)); |
| 220 | if (it == invokers.end()) { |
| 221 | throw std::runtime_error("Adam optimizer with param type "s + c10::toString(params_type) + |
| 222 | " and state type "s + c10::toString(state_type) + |
| 223 | " is not supported on current hardware"s); |
| 224 | } |
| 225 | |
| 226 | it->second(opt, |
| 227 | params.data_ptr(), |
| 228 | grads.data_ptr(), |
| 229 | exp_avg.data_ptr(), |
| 230 | exp_avg_sq.data_ptr(), |
| 231 | param_size, |
| 232 | parallel); |
| 233 | } |
| 234 | |
| 235 | int ds_adam_step(int optimizer_id, |
| 236 | size_t step, |
no test coverage detected