ge 接口 build graph 失败
收藏回复举报
ge 接口 build graph 失败
t('forum.solved') 已解决
发表于2024-03-04 17:03:07
0 查看

使用 ge 接口构图,再调用 aclgrphBuildModel 接口 编译模型失败,代码如下:

#include <cctype>
#include <fstream>
#include <functional>
#include <iostream>
#include <map>
#include <numeric>
#include <string>
#include <unordered_map>
#include <unordered_set>
#include <vector>

#include "acl/acl.h"
#include "all_ops.h"
#include "ascend_string.h"
#include "ge_api.h"
#include "ge_api_types.h"
#include "ge_error_codes.h"
#include "ge_ir_build.h"
#include "gnode.h"
#include "graph.h"
#include "tensor.h"
#include "types.h"

using namespace ge;

class AclgraphBuilder {
 public:
  explicit AclgraphBuilder() {
    // 1. system init
    auto kSocVersion = aclrtGetSocName();
    std::map<AscendString, AscendString> global_options = {
        {AscendString(ge::ir_option::SOC_VERSION), AscendString(kSocVersion)},
    };
    auto status = aclgrphBuildInitialize(global_options);
    if (status != GRAPH_SUCCESS) {
      std::cout << "aclgrphBuildInitialize failed!" << std::endl;
    } else {
      std::cout << "aclgrphBuildInitialize success!" << std::endl;
    }
  }

  void saveGraph(const std::string& path, const Graph& graph) {
    ModelBufferData model;
    std::map<AscendString, AscendString> options;
    auto status = aclgrphBuildModel(graph, options, model);
    if (status == GRAPH_SUCCESS) {
      std::cout << "Build Model SUCCESS!" << std::endl;
    } else {
      std::cout << "Build Model Failed! " << status << std::endl;
      return;
    }

    // 4. Save Ir Model
    status = aclgrphSaveModel(path.c_str(), model);
    if (status == GRAPH_SUCCESS) {
      std::cout << "Save Offline Model SUCCESS!" << std::endl;
    } else {
      std::cout << "Save Offline Model Failed! " << status << std::endl;
    }
  }

  ~AclgraphBuilder() {
    aclgrphBuildFinalize();
    std::cout << "aclgrphBuildFinalize success!" << std::endl;
  }

 private:
  std::string _fusion_switch_file;
};

ge::Operator genInput(const std::string op_name,
                      const std::vector<int64_t> shape, ge::Format format,
                      ge::DataType data_type) {
  TensorDesc tensor_desc_data_op =
      TensorDesc(ge::Shape(shape), format, data_type);
  auto op = op::Data(op_name.c_str());
  op.update_input_desc_x(tensor_desc_data_op);
  op.update_output_desc_y(tensor_desc_data_op);
  return op;
}

void buildGraph(Graph& graph) {

  // input, data nodes
  std::vector<int64_t> input_shape {1, 40, 128};
  auto q = genInput("q", input_shape, FORMAT_ND, ge::DataType::DT_FLOAT16);
  auto k = genInput("k", input_shape, FORMAT_ND, ge::DataType::DT_FLOAT16);
  auto v = genInput("v", input_shape, FORMAT_ND, ge::DataType::DT_FLOAT16);
  graph.AddOp(q);
  graph.AddOp(k);
  graph.AddOp(v);

  // set increflashattention op
  auto incre_fa = op::IncreFlashAttention("increFlashAttention");
  incre_fa.create_dynamic_input_key(1);
  incre_fa.create_dynamic_input_value(1);
  incre_fa.set_input_query(q);
  incre_fa.set_dynamic_input_key(0, k);
  incre_fa.set_dynamic_input_value(0, v);
  incre_fa.SetAttr("num_heads", 1);
  incre_fa.SetAttr("input_layout", "BSH");
  graph.AddOp(incre_fa);
  
  // set input and output
  std::vector<ge::Operator> graph_inputs {q, k, v};
  std::vector<ge::Operator> graph_outputs {incre_fa};
  graph.SetInputs(graph_inputs).SetOutputs(graph_outputs);
}

static void compile() {
  std::string graph_name = "BuildGraph";
  Graph graph(graph_name.c_str());
  buildGraph(graph);

  AclgraphBuilder builder{};
  builder.saveGraph("graph", graph);
  std::cout << "########## end of compile!!!" << std::endl;
}

int main() {
  compile();
  return 0;
}

编译命令:

/usr/bin/c++ -D_GLIBCXX_USE_CXX11_ABI=0 -fPIC -std=c++11 -O3 -Wall -I/usr/local/Ascend/ascend-toolkit/latest/include -I/usr/local/Ascend/ascend-toolkit/latest/opp/built-in/op_proto/inc -I/usr/local/Ascend/ascend-toolkit/latest/include/graph -I/usr/local/Ascend/ascend-toolkit/latest/include/ge -I/usr/local/Ascend/ascend-toolkit/latest/parser -I/usr/local/Ascend/ascend-toolkit/latest/compiler/include -L/usr/local/Ascend/ascend-toolkit/latest/compiler/lib64/stub -lgraph -lge_runner graph_compile.cpp -o./graph_compile /usr/local/Ascend/ascend-toolkit/latest/compiler/lib64/stub/libgraph.so /usr/local/Ascend/ascend-toolkit/latest/compiler/lib64/stub/libge_runner.so /usr/local/Ascend/ascend-toolkit/latest/lib64/libgraph_base.so /usr/local/Ascend/ascend-toolkit/latest/runtime/lib64/stub/libascendcl.so

ascend log: true

主要就是用了一个 IncreFlashAttention 算子 true

代码和错误文件:链接:https://pan.baidu.com/s/1tEqnuG5la-JezcwJF7eHbQ?pwd=xcu0 提取码:xcu0

CANN 版本: 8.0.RC1.alpha001

cpu架构:aarch64

硬件:Atlas 800T A2

我要发帖子