(namespace, model, params, model_checksum,
params_checksum, output, gencode_params)
| 152 | |
| 153 | |
| 154 | def save_model_to_code(namespace, model, params, model_checksum, |
| 155 | params_checksum, output, gencode_params): |
| 156 | util.mkdir_p(output) |
| 157 | cwd = os.path.dirname(__file__) |
| 158 | j2_env = Environment( |
| 159 | loader=FileSystemLoader(cwd + "/template"), trim_blocks=True) |
| 160 | j2_env.filters["stringfy"] = stringfy |
| 161 | |
| 162 | graph_size = len(model.net_def) |
| 163 | for i in range(graph_size): |
| 164 | save_graph_to_code(j2_env, namespace, i, model.net_def[i], output) |
| 165 | |
| 166 | if gencode_params: |
| 167 | template_name = "tensor_data.jinja2" |
| 168 | source = j2_env.get_template(template_name).render( |
| 169 | model_tag=namespace, |
| 170 | model_data_size=len(params), |
| 171 | model_data=params) |
| 172 | with open(output + "/tensor_data.cc", "w") as f: |
| 173 | f.write(source) |
| 174 | |
| 175 | # generate model source files |
| 176 | build_time = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") |
| 177 | template_name = "model.jinja2" |
| 178 | checksum = "{},{}".format(model_checksum, params_checksum) |
| 179 | source = j2_env.get_template(template_name).render( |
| 180 | multi_net=model, |
| 181 | model_tag=namespace, |
| 182 | checksum=checksum, |
| 183 | build_time=build_time) |
| 184 | with open(output + "/model.cc", "w") as f: |
| 185 | f.write(source) |
| 186 | |
| 187 | template_name = 'model_header.jinja2' |
| 188 | source = j2_env.get_template(template_name).render(model_tag=namespace) |
| 189 | with open(output + "/" + namespace + '.h', "w") as f: |
| 190 | f.write(source) |
| 191 | |
| 192 | |
| 193 | def save_model_to_file(model_name, model, params, output): |
no test coverage detected