已合并
fix updateweight bug #789
lilinjie11创建于 4月16日
fix updateweight bug #789
已合并
共 3 个文件变更+116-105
| @@ -16,88 +16,88 @@ | |||
| 16 | Test for MindSpore Lite update_weights | 16 | Test for MindSpore Lite update_weights |
| 17 | """ | 17 | """ |
| 18 | 18 | ||
| 19 | -# import os | 19 | +import os |
| 20 | import pytest | 20 | import pytest |
| 21 | import mindspore_lite as mslite | 21 | import mindspore_lite as mslite |
| 22 | -# import numpy as np | 22 | +import numpy as np |
| 23 | 23 | ||
| 24 | MODEL_FILE = "./single_matmul_model.onnx.mindir" | 24 | MODEL_FILE = "./single_matmul_model.onnx.mindir" |
| 25 | DEVICE_ID = 0 | 25 | DEVICE_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 time | 30 | + 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_ID | 36 | + 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 exc | 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) | ||
| 63 | 46 | ||
| 64 | -# def test_update_weight_zero_copy(): | 47 | +def test_update_weight_multiple_times(): |
| 65 | -# ''' | 48 | + ''' |
| 66 | -# test update weight use zero copy | 49 | + 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_ID | 54 | + 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 weight | 66 | + 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_ID | 71 | + 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-5 | 81 | + 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 | ||
| 102 | def test_update_weight_empty_weight(): | 102 | def 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 model | 117 | + 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_path | 128 | + " --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_ID | 136 | + 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 exc | 141 | + 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 | ||
| 1723 | void DfGraphConvertor::IdentityOptimization() { | 1733 | void 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; |