已合并
fix updateweight bug #789
fix updateweight bug #789
已合并
lilinjie11创建于 4月16日
共 3 个文件变更+116-105
@@ -16,88 +16,88 @@
16Test for MindSpore Lite update_weights16Test for MindSpore Lite update_weights
17"""17"""
18 18 
19-# import os19+import os
20import pytest20import pytest
21import mindspore_lite as mslite21import mindspore_lite as mslite
22-# import numpy as np22+import numpy as np
23 23 
24MODEL_FILE = "./single_matmul_model.onnx.mindir"24MODEL_FILE = "./single_matmul_model.onnx.mindir"
25DEVICE_ID = 025DEVICE_ID = 0
26 26 
27-# TODO: Enable the following commented tests after fix the bug of update weight in CANN.
28-# def test_update_weight_resul_change():
29-# '''
30-# test inference result changed after update weight
31-# '''
32-# model = mslite.Model()
33-# context = mslite.Context()
34-# context.target = ["ascend"]
35-# context.ascend.device_id = DEVICE_ID
36-# model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)
37-# np_input = np.ones((1, 4), dtype=np.float32)
38-# ms_inputs = model.get_inputs()
39-# ms_inputs[0].set_data_from_numpy(np_input)
40-# outputs_nolora = model.predict(ms_inputs)[0].get_data_to_numpy()
41-# weight = np.ones((4, 4), dtype=np.float32)
42-# tensor = mslite.Tensor(weight)
43-# model.update_weights([[tensor]])
44-# outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()
45-# assert not np.allclose(outputs_nolora, outputs_lora)
46 27 
47-# def test_update_weight_multiple_times():28+def test_update_weight_resul_change():
48-# '''29+ '''
49-# test update weight multi time30+ test inference result changed after update weight
50-# '''31+ '''
51-# try:32+ model = mslite.Model()
52-# model = mslite.Model()33+ context = mslite.Context()
53-# context = mslite.Context()34+ context.target = ["ascend"]
54-# context.target = ["ascend"]35+ context.ascend.device_id = DEVICE_ID
55-# context.ascend.device_id = DEVICE_ID36+ model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)
56-# model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)37+ np_input = np.ones((1, 4), dtype=np.float32)
57-# weight = np.ones((4, 4), dtype=np.float32)38+ ms_inputs = model.get_inputs()
58-# tensor = mslite.Tensor(weight)39+ ms_inputs[0].set_data_from_numpy(np_input)
59-# for _ in range(5):40+ outputs_nolora = model.predict(ms_inputs)[0].get_data_to_numpy()
60-# model.update_weights([[tensor]])41+ weight = np.ones((4, 4), dtype=np.float32)
61-# except Exception as exc:42+ tensor = mslite.Tensor(weight)
62-# raise RuntimeError("test update weight multiple times failed!") from exc43+ model.update_weights([[tensor]])
44+ outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()
45+ assert not np.allclose(outputs_nolora, outputs_lora)
63 46 
64-# def test_update_weight_zero_copy():47+def test_update_weight_multiple_times():
65-# '''48+ '''
66-# test update weight use zero copy49+ test update weight multi time
67-# '''50+ '''
68-# model = mslite.Model()51+ try:
69-# context = mslite.Context()52+ model = mslite.Model()
70-# context.target = ["ascend"]53+ context = mslite.Context()
71-# context.ascend.device_id = DEVICE_ID54+ context.target = ["ascend"]
72-# model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)55+ context.ascend.device_id = DEVICE_ID
73-# np_input = np.ones((1, 4), dtype=np.float32)56+ model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)
74-# ms_inputs = model.get_inputs()57+ weight = np.ones((4, 4), dtype=np.float32)
75-# ms_inputs[0].set_data_from_numpy(np_input)58+ tensor = mslite.Tensor(weight)
76-# weight = np.ones((4, 4), dtype=np.float32)59+ for _ in range(5):
77-# tensor = mslite.Tensor(tensor=weight, device="ascend:"+str(DEVICE_ID))60+ model.update_weights([[tensor]])
78-# model.update_weights([[tensor]])61+ except Exception as exc:
79-# outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()62+ raise RuntimeError("test update weight multiple times failed!") from exc
80-# lora_out = np.ones((1, 4), dtype=np.float32) @ np.ones((4, 4), dtype=np.float32)
81-# assert np.mean(lora_out-outputs_lora) < 1e-5
82 63 
83-# def test_update_weight_precision():64+def test_update_weight_zero_copy():
84-# '''65+ '''
85-# test precision after update weight66+ test update weight use zero copy
86-# '''67+ '''
87-# model = mslite.Model()68+ model = mslite.Model()
88-# context = mslite.Context()69+ context = mslite.Context()
89-# context.target = ["ascend"]70+ context.target = ["ascend"]
90-# context.ascend.device_id = DEVICE_ID71+ context.ascend.device_id = DEVICE_ID
91-# model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)72+ model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)
92-# np_input = np.ones((1, 4), dtype=np.float32)73+ np_input = np.ones((1, 4), dtype=np.float32)
93-# ms_inputs = model.get_inputs()74+ ms_inputs = model.get_inputs()
94-# ms_inputs[0].set_data_from_numpy(np_input)75+ ms_inputs[0].set_data_from_numpy(np_input)
95-# weight = np.ones((4, 4), dtype=np.float32)76+ weight = np.ones((4, 4), dtype=np.float32)
96-# tensor = mslite.Tensor(weight)77+ tensor = mslite.Tensor(tensor=weight, device="ascend:"+str(DEVICE_ID))
97-# model.update_weights([[tensor]])78+ model.update_weights([[tensor]])
98-# outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()79+ outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()
99-# lora_out = np.ones((1, 4), dtype=np.float32) @ np.ones((4, 4), dtype=np.float32)80+ lora_out = np.ones((1, 4), dtype=np.float32) @ np.ones((4, 4), dtype=np.float32)
100-# assert np.mean(lora_out-outputs_lora) < 1e-581+ assert np.mean(lora_out-outputs_lora) < 1e-5
82+ 
83+def test_update_weight_precision():
84+ '''
85+ test precision after update weight
86+ '''
87+ model = mslite.Model()
88+ context = mslite.Context()
89+ context.target = ["ascend"]
90+ context.ascend.device_id = DEVICE_ID
91+ model.build_from_file(model_path=MODEL_FILE, model_type=mslite.ModelType.MINDIR, context=context)
92+ np_input = np.ones((1, 4), dtype=np.float32)
93+ ms_inputs = model.get_inputs()
94+ ms_inputs[0].set_data_from_numpy(np_input)
95+ weight = np.ones((4, 4), dtype=np.float32)
96+ tensor = mslite.Tensor(weight)
97+ model.update_weights([[tensor]])
98+ outputs_lora = model.predict(ms_inputs)[0].get_data_to_numpy()
99+ lora_out = np.ones((1, 4), dtype=np.float32) @ np.ones((4, 4), dtype=np.float32)
100+ assert np.mean(lora_out-outputs_lora) < 1e-5
101 101 
102def test_update_weight_empty_weight():102def test_update_weight_empty_weight():
103 '''103 '''
@@ -112,30 +112,30 @@ def test_update_weight_empty_weight():
112 model.update_weights([[]])112 model.update_weights([[]])
113 assert "update weight failed! Error is Common error code" in str(e.value)113 assert "update weight failed! Error is Common error code" in str(e.value)
114 114 
115-# def test_update_weight_mindir(mindir_dir, so_path, output_dir, config_dir):115+def test_update_weight_mindir(mindir_dir, so_path, output_dir, config_dir):
116-# '''116+ '''
117-# test update weight for mindir model117+ test update weight for mindir model
118-# '''118+ '''
119-# model_path = os.path.join(mindir_dir, "linear.mindir")119+ model_path = os.path.join(mindir_dir, "linear.mindir")
120-# fmk_type = "MINDIR"120+ fmk_type = "MINDIR"
121-# config_path = os.path.join(config_dir, "linear.mindir.config")121+ config_path = os.path.join(config_dir, "linear.mindir.config")
122-# output_path = os.path.join(output_dir, "linear_lite")122+ output_path = os.path.join(output_dir, "linear_lite")
123-# cmd_string = so_path + "/tools/converter/converter/converter_lite " + \123+ cmd_string = so_path + "/tools/converter/converter/converter_lite " + \
124-# " --modelFile=" + model_path + \124+ " --modelFile=" + model_path + \
125-# " --optimize=ascend_oriented " + \125+ " --optimize=ascend_oriented " + \
126-# " --outputFile=" + output_path + \126+ " --outputFile=" + output_path + \
127-# " --fmk=" + fmk_type + \127+ " --fmk=" + fmk_type + \
128-# " --configFile=" + config_path128+ " --configFile=" + config_path
129-# ret = os.system(cmd_string)129+ ret = os.system(cmd_string)
130-# if ret != 0:130+ if ret != 0:
131-# raise RuntimeError("model convert failed, cmd_string is: ", cmd_string)131+ raise RuntimeError("model convert failed, cmd_string is: ", cmd_string)
132-# try:132+ try:
133-# model = mslite.Model()133+ model = mslite.Model()
134-# context = mslite.Context()134+ context = mslite.Context()
135-# context.target = ["ascend"]135+ context.target = ["ascend"]
136-# context.ascend.device_id = DEVICE_ID136+ context.ascend.device_id = DEVICE_ID
137-# model.build_from_file(output_path+".mindir", mslite.ModelType.MINDIR, context)137+ model.build_from_file(output_path+".mindir", mslite.ModelType.MINDIR, context)
138-# weight = mslite.Tensor(np.random.randn(64,128).astype(np.float32))138+ weight = mslite.Tensor(np.random.randn(64,128).astype(np.float32))
139-# model.update_weights([[weight]])139+ model.update_weights([[weight]])
140-# except Exception as exc:140+ except Exception as exc:
141-# raise RuntimeError('update weight for mindir model failed!') from exc141+ raise RuntimeError('update weight for mindir model failed!') from exc
@@ -899,7 +899,13 @@ DfGraphConvertor &DfGraphConvertor::BuildGraph(const std::string &name) {
899 MS_LOG(INFO) << "Set graph input num: " << inputs.size();899 MS_LOG(INFO) << "Set graph input num: " << inputs.size();
900 (void)df_graph_->SetInputs(inputs);900 (void)df_graph_->SetInputs(inputs);
901 901 
902- SetGraphOutputs(true);902+ if (IsUpdateGraph()) {
903+ graph_outputs_.clear();
904+ MS_LOG(INFO) << "clear graph outptus";
905+ } else {
906+ SetGraphOutputs(true);
907+ MS_LOG(INFO) << "set graph outptus";
908+ }
903 (void)df_graph_->SetOutputs(graph_outputs_);909 (void)df_graph_->SetOutputs(graph_outputs_);
904 910 
905 IdentityOptimization();911 IdentityOptimization();
@@ -1708,16 +1714,20 @@ void DfGraphConvertor::RemoveIdentity(::ge::GNode identity_node) {
1708 }1714 }
1709}1715}
1710 1716 
1711-bool DfGraphConvertor::IsIdentityInUpdateGraph(const ::ge::GNode &node) const {1717+bool DfGraphConvertor::IsUpdateGraph() const {
1712 MS_EXCEPTION_IF_NULL(anf_graph_);1718 MS_EXCEPTION_IF_NULL(anf_graph_);
1713- auto node_type = GetGNodeType(node);
1714- auto is_identity = (node_type == kTypeIdentityN || node_type == kTypeIdentity);
1715 auto is_update_graph_attr = anf_graph_->get_attr("is_update_graph");1719 auto is_update_graph_attr = anf_graph_->get_attr("is_update_graph");
1716 bool is_update_graph = false;1720 bool is_update_graph = false;
1717 if (is_update_graph_attr != nullptr) {1721 if (is_update_graph_attr != nullptr) {
1718 is_update_graph = GetValue<bool>(is_update_graph_attr);1722 is_update_graph = GetValue<bool>(is_update_graph_attr);
1719 }1723 }
1720- return is_update_graph && is_identity;1724+ return is_update_graph;
1725+}
1726+ 
1727+bool DfGraphConvertor::IsIdentityInUpdateGraph(const ::ge::GNode &node) const {
1728+ auto node_type = GetGNodeType(node);
1729+ auto is_identity = (node_type == kTypeIdentityN || node_type == kTypeIdentity);
1730+ return IsUpdateGraph() && is_identity;
1721}1731}
1722 1732 
1723void DfGraphConvertor::IdentityOptimization() {1733void DfGraphConvertor::IdentityOptimization() {
@@ -234,6 +234,7 @@ class BACKEND_EXPORT DfGraphConvertor {
234 std::string GetGNodeType(const ::ge::GNode &node) const;234 std::string GetGNodeType(const ::ge::GNode &node) const;
235 bool IsIdentityRedundant(const ::ge::GNode &node) const;235 bool IsIdentityRedundant(const ::ge::GNode &node) const;
236 bool IsIdentityInUpdateGraph(const ::ge::GNode &node) const;236 bool IsIdentityInUpdateGraph(const ::ge::GNode &node) const;
237+ bool IsUpdateGraph() const;
237 void RemoveIdentity(::ge::GNode identity_node);238 void RemoveIdentity(::ge::GNode identity_node);
238 void NoOpOptimization();239 void NoOpOptimization();
239 bool IsNoOpRedundant(const ::ge::GNode &node) const;240 bool IsNoOpRedundant(const ::ge::GNode &node) const;