import argparse import json import os import re def _find_metric(metrics, key): if isinstance(metrics, dict): if key in metrics: return metrics[key] for value in metrics.values(): found = _find_metric(value, key) if found is not None: return found if isinstance(metrics, list): for item in metrics: found = _find_metric(item, key) if found is not None: return found return None def _parse_yaml_value(value): value = value.strip() if not value or value in ("null", "~"): return None if value == "true": return True if value == "false": return False if value.startswith("[") and value.endswith("]"): try: return json.loads(value) except (json.JSONDecodeError, ValueError): pass try: return int(value) except ValueError: pass try: return float(value) except ValueError: pass if (value.startswith('"') and value.endswith('"')) or ( value.startswith("'") and value.endswith("'") ): return value[1:-1] return value def _load_params_simple(dirpath): path = os.path.join(dirpath, "params.yaml") alt_path = os.path.join(dirpath, "params-eval.yaml") for p in [path, alt_path]: if os.path.isfile(p): try: with open(p, "r") as f: return _parse_simple_yaml(f.read()) except (OSError, ValueError): continue return None def _parse_simple_yaml(text): lines = text.split("\n") def _effective(raw, indent): stripped = raw.lstrip(" ") cur_indent = len(raw) - len(stripped) if cur_indent != indent: return None, cur_indent if " #" in stripped: effective = stripped[: stripped.index(" #")].rstrip() else: effective = stripped return effective, cur_indent def _peek_nonblank(idx): while idx < len(lines): raw = lines[idx].rstrip() if raw.strip() and not raw.strip().startswith("#"): return idx idx += 1 return len(lines) def _collect_list_items(start, indent): items = [] i = start while i < len(lines): raw = lines[i].rstrip() if not raw.strip() or raw.strip().startswith("#"): i += 1 continue effective, cur_indent = _effective(raw, indent) if effective is None: if cur_indent < indent: break i += 1 continue if effective.startswith("- "): items.append(_parse_yaml_value(effective[2:])) i += 1 else: break return items, i def _parse_block(start, indent): result = {} list_items = [] is_list = False i = start while i < len(lines): raw = lines[i].rstrip() if not raw.strip() or raw.strip().startswith("#"): i += 1 continue stripped = raw.lstrip(" ") cur_indent = len(raw) - len(stripped) if cur_indent < indent: break if cur_indent > indent: i += 1 continue if " #" in stripped: effective = stripped[: stripped.index(" #")].rstrip() else: effective = stripped if not effective: i += 1 continue if effective.startswith("- "): is_list = True list_items.append(_parse_yaml_value(effective[2:])) i += 1 elif effective.endswith(":"): key = effective[:-1].strip() nxt = _peek_nonblank(i + 1) if nxt < len(lines): nxt_raw = lines[nxt].rstrip() nxt_stripped = nxt_raw.lstrip(" ") nxt_indent = len(nxt_raw) - len(nxt_stripped) else: nxt_indent = -1 if nxt_indent > indent: sub_val, i = _parse_block(i + 1, nxt_indent) result[key] = sub_val elif nxt_indent == indent and nxt_stripped.startswith("- "): lst, i = _collect_list_items(i + 1, indent) result[key] = lst else: result[key] = None i += 1 elif ": " in effective: key, _, val_str = effective.partition(": ") result[key.strip()] = _parse_yaml_value(val_str) i += 1 else: i += 1 if is_list: return list_items, i return result, i return _parse_block(0, 0)[0] def _get_in_params(params, path): if params is None: return None keys = path.split(".") current = params for key in keys: if isinstance(current, dict) and key in current: current = current[key] else: return None return current _MODEL_PREFIXES = { "vith": "vit_huge", "vitb": "vit_base", "vitl": "vit_large", "vitt": "vit_tiny", "vits": "vit_small", "vitg": "vit_giant", } def _infer_model_from_run_name(run_name): 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 _get_model_name(dirpath, run_name=None): params = _load_params_simple(dirpath) name = None if params is not 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 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 _get_patch_size(params): if params is None: return "unknown" ps = _get_in_params(params, "mask.patch_size") if ps is None: ps = _get_in_params(params, "meta.patch_size") return ps if ps is not None else "unknown" def _get_crop_size(params): if params is None: return "unknown" cs = _get_in_params(params, "data.crop_size") if cs is None: cs = _get_in_params(params, "meta.crop_size") return cs if cs is not None else "unknown" def _collect_rows(root_dir, metric_key, col_paths): rows = [] for dirpath, _, filenames in os.walk(root_dir): for fname in filenames: if not fname.endswith("_metrics.json"): continue path = os.path.join(dirpath, fname) try: with open(path, "r") as f: metrics = json.load(f) except (OSError, json.JSONDecodeError): continue value = _find_metric(metrics, metric_key) if value is None: continue run_name = os.path.basename(os.path.dirname(path)) model_name = _get_model_name(dirpath, run_name) params = _load_params_simple(dirpath) patch_size = _get_patch_size(params) crop_size = _get_crop_size(params) col_values = [ _get_in_params(params, cp) for cp in col_paths ] rows.append( (float(value), model_name, patch_size, crop_size, col_values, run_name, path) ) return rows def _format_value(val, width=None): if val is None: s = "\u2014" elif isinstance(val, bool): s = str(val) elif isinstance(val, int): s = str(val) elif isinstance(val, float): if abs(val) < 0.001 or abs(val) >= 10000: s = f"{val:.2e}" elif val == int(val): s = f"{val:.1f}" else: s = f"{val:.6f}".rstrip("0").rstrip(".") elif isinstance(val, list): s = str(val) else: s = str(val) if width is not None and len(s) > width: s = s[: width - 1] + "\u2026" return s def _col_width(header, values): w = len(header) for val in values: s = _format_value(val) w = max(w, len(s)) return w _BLOCK = "\u2588" _LIGHT_H = "\u2500" _LIGHT_V = "\u2502" _LIGHT_D = "\u253c" _HEAVY_H = "\u2501" _HEAVY_V = "\u2503" _HEAVY_D = "\u254b" def _print_separator(widths, heavy=False): h = _HEAVY_H if heavy else _LIGHT_H d = _HEAVY_D if heavy else _LIGHT_D parts = [] for i, w in enumerate(widths): if i > 0: parts.append(d) segments = max(w, 1) - 1 if segments <= 0: parts.append(h) else: parts.append(h * segments) line = h.join(parts) if len(parts) > 0 else "" print(f" {line}") def _print_row(values, widths, align_right=None): if align_right is None: align_right = [True] * len(values) parts = [] for i, (val, w) in enumerate(zip(values, widths)): s = _format_value(val, w) if align_right[i]: s = s.rjust(w) else: s = s.ljust(w) if i > 0: parts.append(f" {_LIGHT_V} ") parts.append(s) print(" " + "".join(parts)) def _print_header(full_cols, widths): align = [False] + [True] * (len(full_cols) - 2) + [False] _print_row(full_cols, widths, align) _print_separator(widths) def _print_results(model_type, rows, metric_key, col_headers, top, show_run_name): if not rows: return display_rows = rows[:top] n = len(display_rows) rank_width = max(3, len(str(n))) metric_vals = [r[0] for r in display_rows] metric_width = max( len(metric_key), max(len(_format_value(v)) for v in metric_vals) ) col_widths = [] for i, ch in enumerate(col_headers): col_vals = [r[4][i] for r in display_rows] cw = _col_width(ch, col_vals) col_widths.append(cw) widths = [rank_width, metric_width] + col_widths full_cols = ["#", metric_key] + col_headers if show_run_name: run_name_vals = [r[5] for r in display_rows] run_name_width = max(len("run_name"), max(len(v) for v in run_name_vals)) widths.append(run_name_width) full_cols.append("run_name") title = f"{model_type} \u2014 Top {top} by {metric_key}" sep_len = sum(w + 3 for w in widths) + 2 print(_HEAVY_H * sep_len) print(f" {title}") print(_HEAVY_H * sep_len) _print_header(full_cols, widths) for idx, row in enumerate(display_rows, start=1): metric_str = _format_value(row[0]) col_strs = [_format_value(row[4][i]) for i in range(len(col_headers))] vals = [str(idx), metric_str] + col_strs if show_run_name: vals.append(row[5]) align = [False] + [True] * (len(full_cols) - 2) + [False] _print_row(vals, widths, align) print() def main(): parser = argparse.ArgumentParser( formatter_class=argparse.RawDescriptionHelpFormatter, description="""Summarize grid evaluation results grouped by model / patch size / resolution combination. Examples: %(prog)s --top 10 %(prog)s --metric macro_f1 --cols data.batch_size optimization.lr --top 5 """, ) parser.add_argument( "--root", default="experiment_logs/eval-wilds", help="root folder for eval logs (default: %(default)s)", ) parser.add_argument( "--metric", default="macro_f1", help="metric key to rank by (default: %(default)s)", ) parser.add_argument( "--top", type=int, default=10, help="top N results per model type (default: %(default)s)", ) parser.add_argument( "--cols", nargs="*", default=[], help="extra columns from params.yaml, e.g. data.batch_size optimization.lr", ) parser.add_argument( "--no-run-name", action="store_false", dest="show_run_name", help="hide the run_name column", ) args = parser.parse_args() rows = _collect_rows(args.root, args.metric, args.cols) if not rows: print("No metrics found.") return col_headers = [c.split(".")[-1] for c in args.cols] groups = {} for row in rows: model = row[1] patch = row[2] crop = row[3] key = f"{model} patch={patch} resolution={crop}" groups.setdefault(key, []).append(row) for group_name in sorted(groups.keys()): group = groups[group_name] group.sort(key=lambda r: r[0], reverse=True) _print_results(group_name, group, args.metric, col_headers, args.top, args.show_run_name) if __name__ == "__main__": main()