| 299 | } |
| 300 | |
| 301 | void test_nested(bool check_grad) { |
| 302 | using TwoVar = std::pair<SymbolVar, SymbolVar>; |
| 303 | |
| 304 | static auto make_bisect_pred = [](SymbolVar pred, float thresh) -> TwoVar { |
| 305 | SymbolVar lt, ge; |
| 306 | unpack_vector( |
| 307 | opr::CondExecPred::make( |
| 308 | pred, {pred.make_scalar_dt(thresh)}, |
| 309 | opr::CondExecPred::Mode::PIECEWISE), |
| 310 | lt, ge); |
| 311 | return {lt, ge}; |
| 312 | }; |
| 313 | static auto mark_two = [](SymbolVar x, TwoVar ppvs) -> TwoVar { |
| 314 | SymbolVar a, b; |
| 315 | unpack_vector(opr::CondExecMark::make(ppvs.first, {x}), a); |
| 316 | unpack_vector(opr::CondExecMark::make(ppvs.second, {x}), b); |
| 317 | return {a, b}; |
| 318 | }; |
| 319 | static auto make_bisect = [](SymbolVar x, SymbolVar pred, float thresh, |
| 320 | int* call_lt, int* call_ge, |
| 321 | TwoVar* pred_marked = nullptr) -> TwoVar { |
| 322 | TwoVar pred_br; |
| 323 | SymbolVar x_lt, x_ge; |
| 324 | pred_br = make_bisect_pred(pred, thresh); |
| 325 | std::tie(x_lt, x_ge) = mark_two(x, pred_br); |
| 326 | if (pred_marked) { |
| 327 | *pred_marked = mark_two(pred, pred_br); |
| 328 | } |
| 329 | return {make_call_rec(x_lt, call_lt), make_call_rec(x_ge, call_ge)}; |
| 330 | }; |
| 331 | |
| 332 | auto graph = ComputingGraph::make(); |
| 333 | HostTensorGenerator<> gen; |
| 334 | auto host_x = gen({2, 3}), host_pred = gen({1}); |
| 335 | |
| 336 | int call_lt0, call_ge0; |
| 337 | SymbolVar x = opr::Host2DeviceCopy::make(*graph, host_x).rename("x"), |
| 338 | pred = opr::Host2DeviceCopy::make(*graph, host_pred).rename("pred"), |
| 339 | x_lt_0, x_ge_0; |
| 340 | TwoVar pred_th0; |
| 341 | std::tie(x_lt_0, x_ge_0) = make_bisect(x, pred, 0, &call_lt0, &call_ge0, &pred_th0); |
| 342 | |
| 343 | x_lt_0 = x_lt_0.rename("lt0") / 2; |
| 344 | x_ge_0 = x_ge_0.rename("ge0") * 2; |
| 345 | |
| 346 | int call_n0, call_n1, call_p0, call_p1; |
| 347 | SymbolVar xn0, xn1, xp0, xp1; |
| 348 | std::tie(xn0, xn1) = make_bisect( |
| 349 | x_lt_0, pred_th0.first.rename("pred-neg"), -1, &call_n0, &call_n1); |
| 350 | std::tie(xp0, xp1) = make_bisect( |
| 351 | x_ge_0, pred_th0.second.rename("pred-pos"), 1, &call_p0, &call_p1); |
| 352 | |
| 353 | int call_xn, call_xp; |
| 354 | |
| 355 | auto xn_merge = make_call_rec( |
| 356 | merge_one_out( |
| 357 | {xn0.rename("xn0") - 3, xn1.rename("xn1") + 3}, |
| 358 | MergeMode::EXACT_ONE_SAME_SHAPE), |
no test coverage detected