diff --git a/tools/Polygraphy/CHANGELOG.md b/tools/Polygraphy/CHANGELOG.md index cea3567e..542e3937 100644 --- a/tools/Polygraphy/CHANGELOG.md +++ b/tools/Polygraphy/CHANGELOG.md @@ -3,6 +3,12 @@ Dates are in YYYY-MM-DD format. +## Unreleased +### Fixed +- Parsing a tensor shape no longer accepts a missing, extra, or reversed bracket. + `input:[1,2` was read as shape `[1, 2]`. + + ## v0.53.6 (2026-09-22) ### Added - `Comparator.run()` now accepts a `save_input_blob_path` parameter, exposed on the CLI as diff --git a/tools/Polygraphy/polygraphy/tools/args/util/util.py b/tools/Polygraphy/polygraphy/tools/args/util/util.py index 0ace0371..6c0ca37b 100644 --- a/tools/Polygraphy/polygraphy/tools/args/util/util.py +++ b/tools/Polygraphy/polygraphy/tools/args/util/util.py @@ -38,8 +38,17 @@ def cast(val): """ val = str(val.strip()) - if val.strip("[]") != val: - return [cast(elem) for elem in val.strip("[]").split(",")] + if "[" in val or "]" in val: + well_formed = val.startswith("[") and val.endswith("]") + inner = val[1:-1] if well_formed else "" + if not well_formed or "[" in inner or "]" in inner: + G_LOGGER.critical( + f"Could not parse {val} as a list. " + "Lists must start with '[' and end with ']', with no other brackets." + ) + if not inner.strip(): + return [] + return [cast(elem) for elem in inner.split(",")] try: return int(val) # This fails for float strings like '0.0' diff --git a/tools/Polygraphy/tests/tools/args/util/test_util.py b/tools/Polygraphy/tests/tools/args/util/test_util.py index c43def4c..2a8f6855 100644 --- a/tools/Polygraphy/tests/tools/args/util/test_util.py +++ b/tools/Polygraphy/tests/tools/args/util/test_util.py @@ -71,6 +71,11 @@ def test_parse_shape_with_dim_param_quoted(self, name, quote): meta = args_util.parse_meta(meta_args, includes_dtype=False) assert meta[name].shape == ["batch", 3, 224, 224] + @pytest.mark.parametrize("shape", ["[1,2", "1,2]", "]1,2[", "[1,2]]"]) + def test_unbalanced_shape_is_rejected(self, name, shape): + with pytest.raises(PolygraphyException, match="Could not parse"): + args_util.parse_meta([f"{name}:{shape}"], includes_dtype=False) + class TestRunScript: def test_default_args(self):