已合并
C51修改 #2512
AtomGit-Bot创建于 2022年11月8日
C51修改 #2512
已合并
AtomGit-Bot创建于 2022年11月8日
refs/pull/2512/head合入到master
2 个文件变更+50-46
@@ -26,7 +26,7 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
26 ```26 ```
27 url=https://github.com/ShangtongZhang/DeepRL27 url=https://github.com/ShangtongZhang/DeepRL
28 branch=master28 branch=master
29- commit_id=29+ commit_id=13dd18042414ad112bd0bd383a836d8d739e8acf
30 model_name=C5130 model_name=C51
31 ``` 31 ```
32 32
@@ -86,12 +86,13 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
862. 安装依赖。862. 安装依赖。
87 87 
88 ```88 ```
89+ pip install -r requirements.txt
89 pip install mpi4py90 pip install mpi4py
90 git clone https://github.com/openai/baselines.git91 git clone https://github.com/openai/baselines.git
91 cd baselines92 cd baselines
92 pip install -e .93 pip install -e .
93- pip install -r requirements.txt
94 ```94 ```
95+ >**说明:** pip在线安装requirements.txt中tensorflow==2.6.0仅支持x86架构
95 96 
96## 准备数据集<a name="section183221994411"></a>97## 准备数据集<a name="section183221994411"></a>
97 98 
@@ -108,15 +109,15 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
108 ```109 ```
109 python3.7 c51_preprocess.py c51.model c51.stats dataset/states dataset/actions 1000110 python3.7 c51_preprocess.py c51.model c51.stats dataset/states dataset/actions 1000
110 ```111 ```
111- 参数说明:112+ - 参数说明:
112 113
113- “c51.model”:权重文件。114+ - “c51.model”:权重文件。
114 115 
115- “c51.stats”:模型配置文件。116+ - “c51.stats”:模型配置文件。
116 117 
117- “dataset/states”:stats输出的二进制文件(.bin)所在路径。118+ - “dataset/states”:stats输出的二进制文件(.bin)所在路径。
118 119 
119- “dataset/actions”:action输出的二进制文件(.bin)所在路径。120+ - “dataset/actions”:action输出的二进制文件(.bin)所在路径。
120 121
121 运行成功后生成文件:122 运行成功后生成文件:
122 123
@@ -129,13 +130,13 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
129 ```130 ```
130 python3.7 get_dataset_bin.py dataset/states dataset/bin dataset/out131 python3.7 get_dataset_bin.py dataset/states dataset/bin dataset/out
131 ```132 ```
132- 参数说明:133+ - 参数说明:
133 134
134- “dataset/states”:预处理后的数据文件的相对路径。135+ - “dataset/states”:预处理后的数据文件的相对路径。
135 136 
136- “dataset/bin”:生成的数据集文件保存的路径。137+ - “dataset/bin”:生成的数据集文件保存的路径。
137 138 
138- “dataset/out”:生成的数据集文件格式。139+ - “dataset/out”:bin文件推理后的保存根目录
139 140
140 运行成功后生成文件:141 运行成功后生成文件:
141 142
@@ -198,14 +199,14 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
198 199 
199 - 参数说明:200 - 参数说明:
200 201 
201- - --model:为ONNX模型文件。202+ - --model:为ONNX模型文件。
202- - --framework:5代表ONNX模型。203+ - --framework:5代表ONNX模型。
203- - --output:输出的OM模型。204+ - --output:输出的OM模型。
204- - --input\_format:输入数据的格式。205+ - --input\_format:输入数据的格式。
205- - --input\_shape:输入数据的shape。206+ - --input\_shape:输入数据的shape。
206- - --log:日志级别。207+ - --log:日志级别。
207- - --soc\_version:处理器型号。208+ - --soc\_version:处理器型号。
208- - --insert\_op\_conf=aipp\_resnet34.config: AIPP插入节点,通过config文件配置算子信息,功包括图片色域转换、裁剪、归一化,主要用于处理原图输入数据,常与DVPP配合使用,详见下文数据预处理209+ - --op_select_implmode: 高性模式
209 210 
210 运行成功后生成c51_bs1.om模型文件。211 运行成功后生成c51_bs1.om模型文件。
211 212 
@@ -216,20 +217,24 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
216 ais-infer工具获取及使用方式请点击查看[[ais_infer 推理工具使用文档](https://gitee.com/ascend/tools/tree/master/ais-bench_workload/tool/ais_infer)]217 ais-infer工具获取及使用方式请点击查看[[ais_infer 推理工具使用文档](https://gitee.com/ascend/tools/tree/master/ais-bench_workload/tool/ais_infer)]
217 218 
218 b. 执行推理。219 b. 执行推理。
220+ ```shell
221+ python3.7 ais_infer.py --model=c51_bs1.om --input dataset/bin --output dataset/out/2022_11_8_21_03_50 --outfmt TXT --batchsize 1
222+ ```
219 223 
220- ` python3.7 ${ais_infer_path}/ais_infer.py --model=${om_model_path} --loop=20 --batchsize=${batch_size} `
221 224 
222- - ${om_path}: 之前生成的OM模型的位置225+ - 参数说明:
223- 
224- - ${Bin_data_path}: 数据预处理后,二进制文件所在目录
225 226
226- - --model: 需要进行推理的om模型227+ - --model: om模型的路径
227- 228+ 
228- - --output: 推理结果出路径。229+ - --input: 入的bin文件目录
229- 230+
230- - --outfmt: 输出数据的格式,默认”BIN“,可取值“NPY”、“BIN”、“TXT”231+ - --output: 推理结果输出路径
231- 232+
232- - --input: 模型需要的入,支持bin文件和目录,若不加该参,会自动生成都为0的数233+ - --outfmt: 数据的格式
234+ 
235+ - --batchsize : 模型输入批次大小
236+
237+ 
233 238
234 说明: 执行ais-infer工具请选择与运行环境架构相同的命令。239 说明: 执行ais-infer工具请选择与运行环境架构相同的命令。
235 240 
@@ -238,14 +243,14 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
238 调用脚本与数据集标签比对,可以获得Accuracy数据。243 调用脚本与数据集标签比对,可以获得Accuracy数据。
239 244 
240 ```245 ```
241- python3.7 c51_postprocess.py dataset/actions dataset/out 1000246+ python3.7 c51_postprocess.py dataset/actions dataset/out/2022_11_8_21_03_50 1000
242 ```247 ```
243 248 
244- - 参数说明:249+ - 参数说明:
245 250 
246 - “dataset/actions”:保存的输出action的路径。251 - “dataset/actions”:保存的输出action的路径。
247 252 
248- - “dataset/out”:离线推理输出的路径。253+ - “dataset/out/2022_11_8_21_03_50”:离线推理输出的路径。
249 254 
250 - “1000”:参数输出比较的个数。255 - “1000”:参数输出比较的个数。
251 256 
@@ -256,9 +261,9 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
256 261 
257调用ACL接口推理计算,性能参考下列数据。262调用ACL接口推理计算,性能参考下列数据。
258 263 
259-| batch_size | 310 | 310P | T4 | 310P/310 | 310P/T4 |264+| batch_size | 310P |
260-|------------|----------|---------|----------|----------|---------|265+|------------|---------|
261-| bs1 | 13572.84 | 6050.12 | 15574.14 | 0.44575 | 0.38847|266+| bs1 | 6050.12 |
262 267 
263精度参考下列数据。268精度参考下列数据。
264 269 
@@ -267,3 +272,4 @@ C51是一种值分布强化学习算法,C51算法的框架依然是DQN算法
267| 310P精度 | 98.9% |272| 310P精度 | 98.9% |
268 273 
269注:此模型不支持多batch。274注:此模型不支持多batch。
275+
@@ -48,20 +48,18 @@ def get_pth_action(filename):
48 48 
49 49 
50if __name__ == "__main__":50if __name__ == "__main__":
51- action_file = sys.argv[1]51+ action_file = sys.argv[1] # dataset/action
52- out_file = sys.argv[2]52+ out_file = sys.argv[2] # dataset/out
53- num = int(sys.argv[3])53+ num = int(sys.argv[3]) # 1000
54- out_dir = os.listdir(out_file)
55- for out_dir_file in out_dir:
56- om_filelist = os.listdir('{0}/{1}'.format(out_file, out_dir_file))
57- file_num = len(om_filelist)
58 equal = 054 equal = 0
59- for i in range(file_num):55+ om_filelist = []
56+ for i, om_result_path in enumerate(os.listdir(out_file)):
57+ om_filelist.append(om_result_path)
60 pth_action = get_pth_action('{0}/{1}.pt'.format(action_file, i))58 pth_action = get_pth_action('{0}/{1}.pt'.format(action_file, i))
61- om_action = get_om_action('{0}/{1}/{2}_output_0.txt'.format(out_file, out_dir_file, i))59+ om_action = get_om_action('{0}/{1}_0.txt'.format(out_file, i))
62 if pth_action==om_action:60 if pth_action==om_action:
63 equal += 161 equal += 1
64- print('om离线推理的精度是在线推理的:{0}'.format((equal/num)))62+ print('The offline inference accuracy of om is {0} times higher than the online inference accuracy'.format((equal/num)))
65 if(equal > 0.9*num):63 if(equal > 0.9*num):
66 print("Accuancy: OK")64 print("Accuancy: OK")
67 else:65 else: