Unary dst is the output dtype of op_node.
| 41 | // Unary |
| 42 | // dst is the output dtype of op_node. |
| 43 | Status Unary(const FDH::Node& op_node, const Tensor& x, const DataType dst, |
| 44 | Tensor* y) { |
| 45 | const DataType src = x.dtype(); |
| 46 | auto adef = [](const string& name, |
| 47 | const DataType type) { // E.g., x:float, dy:double |
| 48 | return strings::StrCat(name, ":", DataTypeString(type)); |
| 49 | }; |
| 50 | // Sum(op(x)), sum all output of op(x). |
| 51 | auto test = FDH::Define("Test", {adef("x", src)}, {adef("l", dst)}, {}, |
| 52 | { |
| 53 | op_node, |
| 54 | FDH::Const("zero", 0), |
| 55 | FDH::Const("one", 1), |
| 56 | {{"r"}, "Rank", {"x"}, {{"T", src}}}, |
| 57 | {{"indices"}, "Range", {"zero", "r", "one"}}, |
| 58 | {{"l"}, "Sum", {"y", "indices"}, {{"T", dst}}}, |
| 59 | }); |
| 60 | |
| 61 | // TestGrad = Test'(x) |
| 62 | auto grad = FDH::Define( |
| 63 | "TestGrad", {adef("x", src)}, {adef("dx", src)}, {}, |
| 64 | { |
| 65 | FDH::Const("one", 1), |
| 66 | {{"dy"}, "Cast", {"one"}, {{"DstT", dst}, {"SrcT", DT_INT32}}}, |
| 67 | {{"grad"}, |
| 68 | "SymbolicGradient", |
| 69 | {"x", "dy"}, |
| 70 | { |
| 71 | {"f", FDH::FunctionRef("Test")}, |
| 72 | {"Tin", DataTypeSlice{src, dst}}, |
| 73 | {"Tout", DataTypeSlice{src}}, |
| 74 | }}, |
| 75 | {{"dx"}, "Identity", {"grad"}, {{"T", src}}}, |
| 76 | }); |
| 77 | // Each test case will feed in "x:0" and expects to get "dx:0". |
| 78 | auto gdef = test::function::GDef( |
| 79 | { |
| 80 | f::NDef("x", "Placeholder", {}, {{"dtype", src}}), |
| 81 | f::NDef("dx", "TestGrad", {"x"}, {}), |
| 82 | }, |
| 83 | {test, grad}); |
| 84 | |
| 85 | auto sess = NewSession(); |
| 86 | TF_CHECK_OK(sess->Create(gdef)); |
| 87 | std::vector<Tensor> outputs; |
| 88 | auto s = sess->Run({{"x:0", x}}, {"dx:0"}, {}, &outputs); |
| 89 | if (s.ok()) { |
| 90 | CHECK_EQ(outputs.size(), 1); |
| 91 | *y = outputs[0]; |
| 92 | } |
| 93 | TF_CHECK_OK(sess->Close()); |
| 94 | return s; |
| 95 | } |
| 96 | |
| 97 | Status Unary(const string& op, const Tensor& x, Tensor* y) { |
| 98 | const FDH::Node op_node = {{"y"}, op, {"x"}, {{"T", x.dtype()}}}; |
nothing calls this directly
no test coverage detected