patch size and resolution filter
This commit is contained in:
+42
-10
@@ -228,6 +228,24 @@ def _get_model_name(dirpath, run_name=None):
|
|||||||
return base
|
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):
|
def _collect_rows(root_dir, metric_key, col_paths):
|
||||||
rows = []
|
rows = []
|
||||||
for dirpath, _, filenames in os.walk(root_dir):
|
for dirpath, _, filenames in os.walk(root_dir):
|
||||||
@@ -246,11 +264,13 @@ def _collect_rows(root_dir, metric_key, col_paths):
|
|||||||
run_name = os.path.basename(os.path.dirname(path))
|
run_name = os.path.basename(os.path.dirname(path))
|
||||||
model_name = _get_model_name(dirpath, run_name)
|
model_name = _get_model_name(dirpath, run_name)
|
||||||
params = _load_params_simple(dirpath)
|
params = _load_params_simple(dirpath)
|
||||||
|
patch_size = _get_patch_size(params)
|
||||||
|
crop_size = _get_crop_size(params)
|
||||||
col_values = [
|
col_values = [
|
||||||
_get_in_params(params, cp) for cp in col_paths
|
_get_in_params(params, cp) for cp in col_paths
|
||||||
]
|
]
|
||||||
rows.append(
|
rows.append(
|
||||||
(float(value), model_name, col_values, run_name, path)
|
(float(value), model_name, patch_size, crop_size, col_values, run_name, path)
|
||||||
)
|
)
|
||||||
return rows
|
return rows
|
||||||
|
|
||||||
@@ -348,7 +368,7 @@ def _print_results(model_type, rows, metric_key, col_headers, top, show_run_name
|
|||||||
|
|
||||||
col_widths = []
|
col_widths = []
|
||||||
for i, ch in enumerate(col_headers):
|
for i, ch in enumerate(col_headers):
|
||||||
col_vals = [r[2][i] for r in display_rows]
|
col_vals = [r[4][i] for r in display_rows]
|
||||||
cw = _col_width(ch, col_vals)
|
cw = _col_width(ch, col_vals)
|
||||||
col_widths.append(cw)
|
col_widths.append(cw)
|
||||||
|
|
||||||
@@ -356,7 +376,7 @@ def _print_results(model_type, rows, metric_key, col_headers, top, show_run_name
|
|||||||
full_cols = ["#", metric_key] + col_headers
|
full_cols = ["#", metric_key] + col_headers
|
||||||
|
|
||||||
if show_run_name:
|
if show_run_name:
|
||||||
run_name_vals = [r[3] for r in display_rows]
|
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))
|
run_name_width = max(len("run_name"), max(len(v) for v in run_name_vals))
|
||||||
widths.append(run_name_width)
|
widths.append(run_name_width)
|
||||||
full_cols.append("run_name")
|
full_cols.append("run_name")
|
||||||
@@ -370,10 +390,10 @@ def _print_results(model_type, rows, metric_key, col_headers, top, show_run_name
|
|||||||
|
|
||||||
for idx, row in enumerate(display_rows, start=1):
|
for idx, row in enumerate(display_rows, start=1):
|
||||||
metric_str = _format_value(row[0])
|
metric_str = _format_value(row[0])
|
||||||
col_strs = [_format_value(row[2][i]) for i in range(len(col_headers))]
|
col_strs = [_format_value(row[4][i]) for i in range(len(col_headers))]
|
||||||
vals = [str(idx), metric_str] + col_strs
|
vals = [str(idx), metric_str] + col_strs
|
||||||
if show_run_name:
|
if show_run_name:
|
||||||
vals.append(row[3])
|
vals.append(row[5])
|
||||||
align = [False] + [True] * (len(full_cols) - 2) + [False]
|
align = [False] + [True] * (len(full_cols) - 2) + [False]
|
||||||
_print_row(vals, widths, align)
|
_print_row(vals, widths, align)
|
||||||
|
|
||||||
@@ -418,6 +438,12 @@ Examples:
|
|||||||
dest="show_run_name",
|
dest="show_run_name",
|
||||||
help="hide the run_name column",
|
help="hide the run_name column",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--group-by",
|
||||||
|
default="model",
|
||||||
|
choices=["model", "patch_size", "crop_size"],
|
||||||
|
help="group results by this attribute (default: %(default)s)",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
rows = _collect_rows(args.root, args.metric, args.cols)
|
rows = _collect_rows(args.root, args.metric, args.cols)
|
||||||
@@ -428,15 +454,21 @@ Examples:
|
|||||||
|
|
||||||
col_headers = [c.split(".")[-1] for c in args.cols]
|
col_headers = [c.split(".")[-1] for c in args.cols]
|
||||||
|
|
||||||
|
group_key = {
|
||||||
|
"model": lambda r: r[1],
|
||||||
|
"patch_size": lambda r: str(r[2]),
|
||||||
|
"crop_size": lambda r: str(r[3]),
|
||||||
|
}[args.group_by]
|
||||||
|
|
||||||
groups = {}
|
groups = {}
|
||||||
for row in rows:
|
for row in rows:
|
||||||
model = row[1]
|
key = group_key(row)
|
||||||
groups.setdefault(model, []).append(row)
|
groups.setdefault(key, []).append(row)
|
||||||
|
|
||||||
for model_type in sorted(groups.keys()):
|
for group_name in sorted(groups.keys()):
|
||||||
group = groups[model_type]
|
group = groups[group_name]
|
||||||
group.sort(key=lambda r: r[0], reverse=True)
|
group.sort(key=lambda r: r[0], reverse=True)
|
||||||
_print_results(model_type, group, args.metric, col_headers, args.top, args.show_run_name)
|
_print_results(group_name, group, args.metric, col_headers, args.top, args.show_run_name)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user