handle imagenet models
This commit is contained in:
+12
-1
@@ -200,6 +200,15 @@ def _infer_model_from_run_name(run_name):
|
|||||||
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
|
||||||
@@ -207,7 +216,9 @@ 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)
|
||||||
return name or "unknown"
|
if name:
|
||||||
|
return _normalize_model_type(name)
|
||||||
|
return "unknown"
|
||||||
|
|
||||||
|
|
||||||
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