fix: check for run name
This commit is contained in:
+20
-13
@@ -1,6 +1,7 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
|
|
||||||
|
|
||||||
def _find_metric(metrics, key):
|
def _find_metric(metrics, key):
|
||||||
@@ -194,21 +195,16 @@ _MODEL_PREFIXES = {
|
|||||||
|
|
||||||
|
|
||||||
def _infer_model_from_run_name(run_name):
|
def _infer_model_from_run_name(run_name):
|
||||||
for prefix, model in _MODEL_PREFIXES.items():
|
lower = run_name.lower()
|
||||||
if run_name.lower().startswith(prefix):
|
for model in sorted(set(_MODEL_PREFIXES.values()), key=len, reverse=True):
|
||||||
|
if lower.startswith(model):
|
||||||
|
return model
|
||||||
|
for prefix, model in sorted(_MODEL_PREFIXES.items(), key=lambda x: len(x[0]), reverse=True):
|
||||||
|
if lower.startswith(prefix):
|
||||||
return model
|
return model
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _normalize_model_type(model_name):
|
|
||||||
parts = model_name.split("-")
|
|
||||||
base = parts[0]
|
|
||||||
for part in parts[1:]:
|
|
||||||
if part.startswith("in"):
|
|
||||||
return f"{base}-{part}"
|
|
||||||
return base
|
|
||||||
|
|
||||||
|
|
||||||
def _get_model_name(dirpath, run_name=None):
|
def _get_model_name(dirpath, run_name=None):
|
||||||
params = _load_params_simple(dirpath)
|
params = _load_params_simple(dirpath)
|
||||||
name = None
|
name = None
|
||||||
@@ -216,9 +212,20 @@ def _get_model_name(dirpath, run_name=None):
|
|||||||
name = _get_in_params(params, "meta.model_name")
|
name = _get_in_params(params, "meta.model_name")
|
||||||
if name is None and run_name:
|
if name is None and run_name:
|
||||||
name = _infer_model_from_run_name(run_name)
|
name = _infer_model_from_run_name(run_name)
|
||||||
if name:
|
if name is None:
|
||||||
return _normalize_model_type(name)
|
|
||||||
return "unknown"
|
return "unknown"
|
||||||
|
base = name.split("-")[0]
|
||||||
|
for part in name.split("-")[1:]:
|
||||||
|
if part.startswith("in"):
|
||||||
|
return f"{base}-{part}"
|
||||||
|
if run_name:
|
||||||
|
for part in run_name.split("-"):
|
||||||
|
if part.lower().startswith("in"):
|
||||||
|
return f"{base}-{part}"
|
||||||
|
m = re.search(r"in\d+\w*", part)
|
||||||
|
if m:
|
||||||
|
return f"{base}-{m.group()}"
|
||||||
|
return base
|
||||||
|
|
||||||
|
|
||||||
def _collect_rows(root_dir, metric_key, col_paths):
|
def _collect_rows(root_dir, metric_key, col_paths):
|
||||||
|
|||||||
Reference in New Issue
Block a user