fix: check for run name

This commit is contained in:
Yann Ahlgrim
2026-06-20 12:04:30 +02:00
parent 7d36feec44
commit 164e25fe55
+20 -13
View File
@@ -1,6 +1,7 @@
import argparse
import json
import os
import re
def _find_metric(metrics, key):
@@ -194,21 +195,16 @@ _MODEL_PREFIXES = {
def _infer_model_from_run_name(run_name):
for prefix, model in _MODEL_PREFIXES.items():
if run_name.lower().startswith(prefix):
lower = run_name.lower()
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 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):
params = _load_params_simple(dirpath)
name = None
@@ -216,9 +212,20 @@ def _get_model_name(dirpath, run_name=None):
name = _get_in_params(params, "meta.model_name")
if name is None and run_name:
name = _infer_model_from_run_name(run_name)
if name:
return _normalize_model_type(name)
if name is None:
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):