import numpy as np
from ge.es.all import MatMul
from ge.es.graph_builder import GraphBuilder
from ge.ge_global import GeApi
from ge.graph import DumpFormat, Tensor
from ge.graph.types import DataType, Format
from ge.session import Session
def build_matmul_graph():
builder = GraphBuilder("MakeMatMulGraph")
input_tensor_holder = builder.create_input(index=0, name="input", data_type=DataType.DT_FLOAT, shape=[2, 3])
weight = builder.create_const_float([1.0] * 6, shape=[2, 3])
matmul_tensor_holder = MatMul(weight, input_tensor_holder, None, transpose_x1=True, transpose_x2=False)
builder.set_graph_output(matmul_tensor_holder, 0)
return builder.build_and_reset()
def dump_matmul_graph(graph):
graph.dump_to_file(format=DumpFormat.kOnnx, suffix="make_matmul_graph")
def run_graph(graph, device_id="0") -> int:
config = {"ge.exec.deviceId": str(device_id), "ge.graphRunMode": "0"}
ge_api = GeApi()
ret = ge_api.ge_initialize(config)
if ret != 0:
print(f"[Error] GE初始化失败,返回码: {ret}")
return ret
print(f"[Info] GE环境初始化成功 (Device ID: {device_id})")
try:
session = Session()
graph_id = 1
ret = session.add_graph(graph_id, graph)
if ret != 0:
print(f"[Error] 添加图失败,返回码: {ret}")
return ret
print(f"[Info] 图已添加到Session (Graph ID: {graph_id})")
input_data = np.array([[1, 2, 3], [4, 5, 6]], dtype=np.float32)
input_tensor = Tensor(
input_data.flatten().tolist(),
None,
DataType.DT_FLOAT,
Format.FORMAT_ND,
[2, 3],
)
inputs = [input_tensor]
print(f"[Info] 输入数据已准备,共{len(inputs)}个输入tensor")
ret = session.run_graph(graph_id, inputs)
print("[Info] 图运行成功!")
for idx, tensor in enumerate(ret, start=1):
print(f"Tensor{idx}详情:{tensor}")
return 0
except Exception as e:
print(f"[Error] 执行过程中出错: {e}")
import traceback
traceback.print_exc()
return -1
finally:
print("[Info] 清理GE环境...")
ge_api.ge_finalize()
print("[Success] GE环境已清理")
if __name__ == "__main__":
graph = build_matmul_graph()
dump_matmul_graph(graph)
run_graph(graph)