| 51 | MGB_DYN_TYPE_OBJ_FINAL_IMPL(WorkspaceLimitGetterOpr); |
| 52 | |
| 53 | void run_test(bool dynamic) { |
| 54 | HostTensorGenerator<> gen; |
| 55 | auto graph = ComputingGraph::make(); |
| 56 | |
| 57 | if (dynamic) { |
| 58 | graph->options().force_dynamic_alloc = true; |
| 59 | } |
| 60 | |
| 61 | auto x = opr::SharedDeviceTensor::make(*graph, *gen({23})); |
| 62 | |
| 63 | int infer_shape_nr_call = 0; |
| 64 | auto infer_shape_callback = [&]() { |
| 65 | ++infer_shape_nr_call; |
| 66 | if (infer_shape_nr_call < 3) { |
| 67 | ASSERT_TRUE(WorkspaceLimitGetter::is_prealloc_run(graph.get())); |
| 68 | } else { |
| 69 | ASSERT_FALSE(WorkspaceLimitGetter::is_prealloc_run(graph.get())); |
| 70 | auto wk = WorkspaceLimitGetter::get_workspace_limit( |
| 71 | graph.get(), x.node()->comp_node(), 123); |
| 72 | ASSERT_GT(wk, 0u); |
| 73 | ASSERT_LE(wk, 123u); |
| 74 | return; |
| 75 | } |
| 76 | }; |
| 77 | |
| 78 | auto y = WorkspaceLimitGetterOpr::make(x, infer_shape_callback); |
| 79 | ASSERT_EQ(1, infer_shape_nr_call); |
| 80 | |
| 81 | graph->compile({{x, {}}})->execute(); |
| 82 | ASSERT_EQ(1, infer_shape_nr_call); |
| 83 | |
| 84 | auto func1 = graph->compile({{y, {}}}); |
| 85 | ASSERT_EQ(1, infer_shape_nr_call); |
| 86 | func1->execute(); |
| 87 | ASSERT_EQ(3, infer_shape_nr_call); |
| 88 | |
| 89 | func1->execute(); |
| 90 | ASSERT_EQ(3, infer_shape_nr_call); |
| 91 | } |
| 92 | |
| 93 | } // namespace |
| 94 | |