| 416 | #undef REGISTER_OP_BY_NAME |
| 417 | |
| 418 | static Status ApplyAdamAsyncShapeFn(InferenceContext* c, bool sparse) { |
| 419 | ShapeHandle unused; |
| 420 | ShapeHandle s = ShapeOrHandleShape(c, 0); // var |
| 421 | TF_RETURN_IF_ERROR(c->Merge(s, ShapeOrHandleShape(c, 1), &s)); // m |
| 422 | TF_RETURN_IF_ERROR(c->Merge(s, ShapeOrHandleShape(c, 2), &s)); // v |
| 423 | TF_RETURN_IF_ERROR(c->WithRank(c->input(3), 0, &unused)); // beta1_power |
| 424 | TF_RETURN_IF_ERROR(c->WithRank(c->input(4), 0, &unused)); // beta2_power |
| 425 | TF_RETURN_IF_ERROR(c->WithRank(c->input(5), 0, &unused)); // lr |
| 426 | TF_RETURN_IF_ERROR(c->WithRank(c->input(6), 0, &unused)); // beta1 |
| 427 | TF_RETURN_IF_ERROR(c->WithRank(c->input(7), 0, &unused)); // beta2 |
| 428 | TF_RETURN_IF_ERROR(c->WithRank(c->input(8), 0, &unused)); // epsilon |
| 429 | TF_RETURN_IF_ERROR( |
| 430 | HandleGradAndIndicesInputs(c, sparse, 9 /* grad_idx */, &s)); |
| 431 | if (c->num_outputs() > 0) { |
| 432 | c->set_output(0, s); |
| 433 | } |
| 434 | return Status::OK(); |
| 435 | } |
| 436 | |
| 437 | REGISTER_OP("ApplyAdamAsync") |
| 438 | .Input("var: Ref(T)") |
no test coverage detected