| 656 | |
| 657 | |
| 658 | bool test_graph_save_load() { |
| 659 | try { |
| 660 | const std::string filename = "test_graph_save_load.cg"; |
| 661 | |
| 662 | CactusGraph graph; |
| 663 | size_t input_a = graph.input({2, 3}, Precision::FP16); |
| 664 | size_t input_b = graph.input({2, 3}, Precision::FP16); |
| 665 | size_t sum_id = graph.add(input_a, input_b); |
| 666 | graph.pow(sum_id, 2.0f); |
| 667 | |
| 668 | graph.save(filename); |
| 669 | |
| 670 | GraphFile::SerializedGraph sg = GraphFile::load_graph(filename); |
| 671 | |
| 672 | if (sg.header.node_count != 4) { |
| 673 | std::cout << "[graph_save_load] unexpected node_count: " |
| 674 | << sg.header.node_count << " expected 4" << std::endl; |
| 675 | std::remove(filename.c_str()); |
| 676 | return false; |
| 677 | } |
| 678 | |
| 679 | if (sg.graph_inputs.size() != 2 || sg.graph_inputs[0] != 0 || |
| 680 | sg.graph_inputs[1] != 1) { |
| 681 | std::cout << "[graph_save_load] unexpected graph_inputs:"; |
| 682 | for (uint32_t idx : sg.graph_inputs) { |
| 683 | std::cout << " " << idx; |
| 684 | } |
| 685 | std::cout << std::endl; |
| 686 | std::remove(filename.c_str()); |
| 687 | return false; |
| 688 | } |
| 689 | |
| 690 | if (sg.graph_outputs.size() != 1 || sg.graph_outputs[0] != 3) { |
| 691 | std::cout << "[graph_save_load] unexpected graph_outputs:"; |
| 692 | for (uint32_t idx : sg.graph_outputs) { |
| 693 | std::cout << " " << idx; |
| 694 | } |
| 695 | std::cout << std::endl; |
| 696 | std::remove(filename.c_str()); |
| 697 | return false; |
| 698 | } |
| 699 | |
| 700 | CactusGraph loaded = CactusGraph::load(filename); |
| 701 | |
| 702 | std::vector<__fp16> data_a = {1, 2, 3, 4, 5, 6}; |
| 703 | std::vector<__fp16> data_b = {10, 20, 30, 40, 50, 60}; |
| 704 | |
| 705 | const size_t loaded_input_a = runtime_id_from_serialized_index(sg.graph_inputs[0]); |
| 706 | const size_t loaded_input_b = runtime_id_from_serialized_index(sg.graph_inputs[1]); |
| 707 | const size_t loaded_output_id = runtime_id_from_serialized_index(sg.graph_outputs[0]); |
| 708 | |
| 709 | loaded.set_input(loaded_input_a, data_a.data(), Precision::FP16); |
| 710 | loaded.set_input(loaded_input_b, data_b.data(), Precision::FP16); |
| 711 | loaded.execute(); |
| 712 | |
| 713 | __fp16* output = static_cast<__fp16*>(loaded.get_output(loaded_output_id)); |
| 714 | std::vector<float> expected = { |
| 715 | 121.0f, 484.0f, 1089.0f, |
no test coverage detected