| 7 | from seq2struct.utils import registry |
| 8 | |
| 9 | def compute_metrics(config_path, config_args, section, inferred_path,logdir=None): |
| 10 | if config_args: |
| 11 | config = json.loads(_jsonnet.evaluate_file(config_path, tla_codes={'args': config_args})) |
| 12 | else: |
| 13 | config = json.loads(_jsonnet.evaluate_file(config_path)) |
| 14 | |
| 15 | if 'model_name' in config and logdir: |
| 16 | logdir = os.path.join(logdir, config['model_name']) |
| 17 | if logdir: |
| 18 | inferred_path = inferred_path.replace('__LOGDIR__', logdir) |
| 19 | |
| 20 | inferred = open(inferred_path) |
| 21 | data = registry.construct('dataset', config['data'][section]) |
| 22 | metrics = data.Metrics(data) |
| 23 | |
| 24 | inferred_lines = list(inferred) |
| 25 | if len(inferred_lines) < len(data): |
| 26 | raise Exception('Not enough inferred: {} vs {}'.format(len(inferred_lines), |
| 27 | len(data))) |
| 28 | |
| 29 | |
| 30 | for line in inferred_lines: |
| 31 | infer_results = json.loads(line) |
| 32 | if infer_results['beams']: |
| 33 | inferred_code = infer_results['beams'][0]['inferred_code'] |
| 34 | else: |
| 35 | inferred_code = None |
| 36 | if 'index' in infer_results: |
| 37 | metrics.add(data[infer_results['index']], inferred_code) |
| 38 | else: |
| 39 | metrics.add(None, inferred_code, obsolete_gold_code=infer_results['gold_code']) |
| 40 | |
| 41 | return logdir, metrics.finalize() |