diff --git a/penzai/core/shapecheck.py b/penzai/core/shapecheck.py index 78580ea..4992436 100644 --- a/penzai/core/shapecheck.py +++ b/penzai/core/shapecheck.py @@ -93,7 +93,7 @@ class DimVar(Mapping): and the inner name is the named shape of a single axis in that collection. """ - name: str | tuple[str, str | int] + name: str | tuple[str, named_axes.AxisName | int] def __len__(self) -> int: return 1 @@ -148,7 +148,7 @@ class KnownDim: from_keypath: Optional keypath that indicates where this size was bound. """ - name: str | tuple[str, str | int] + name: str | tuple[str, named_axes.AxisName | int] size: int from_keypath: str | None = None @@ -177,8 +177,11 @@ def var(name: str) -> DimVar: def vars_for_axes( var_name: str, - axis_names_or_specs: Collection[str] | Mapping[str, int | None], -) -> dict[str, DimVar | KnownDim]: + axis_names_or_specs: ( + Collection[named_axes.AxisName] + | Mapping[named_axes.AxisName, int | None] + ), +) -> dict[named_axes.AxisName, DimVar | KnownDim]: """Creates variables for a known collection of named axes. Args: @@ -512,7 +515,7 @@ def _try_match_one( keypath, pattern: int | DimVar | MultiDimVar | KnownDim, value: int | tuple[int, ...] | dict[named_axes.AxisName, int], - solutions: dict[str | tuple[str, str], _Binding], + solutions: dict[str | tuple[str, named_axes.AxisName | int], _Binding], ) -> str | None: """Internal helper to match a pattern with a value. @@ -574,7 +577,7 @@ def _try_match_one( def _positional_inline_multidimvars( constraint: _PositionalConstraint, - solutions: dict[str | tuple[str, str], _Binding], + solutions: dict[str | tuple[str, named_axes.AxisName | int], _Binding], ) -> tuple[_PositionalConstraint, str]: """Simplifies a positional constraint by inlining multivars.""" new_pattern = [] @@ -601,7 +604,7 @@ def _positional_inline_multidimvars( def _named_inline_multidimvars( constraint: _NamedConstraint, - solutions: dict[str | tuple[str, str], _Binding], + solutions: dict[str | tuple[str, named_axes.AxisName | int], _Binding], ) -> tuple[_NamedConstraint | _UnsatisfiedConstraint, str]: """Simplifies a named constraint by inlining multivars.""" new_pattern = {} @@ -618,7 +621,6 @@ def _named_inline_multidimvars( binding = solutions[key.name] assert isinstance(binding.value, dict) for subkey, subval in binding.value.items(): - assert isinstance(subkey, str) if subkey in new_pattern: return ( _UnsatisfiedConstraint( @@ -1038,7 +1040,6 @@ def add_constraints(keypath, pattern: Any, value: Any): ): found = solutions[name[0]].value[name[1]] else: - assert isinstance(name[1], str) if isinstance(solutions[name[0]].value, dict): found = solutions[name[0]].value.get(name[1]) if found != binding.value: diff --git a/tests/core/shapecheck_test.py b/tests/core/shapecheck_test.py index 80afc41..c063624 100644 --- a/tests/core/shapecheck_test.py +++ b/tests/core/shapecheck_test.py @@ -94,9 +94,11 @@ def test_bad_shapes_dtypes(self): def test_simple_named(self): match = pz.chk.check_structure( value={ - "a": pz.nx.zeros( - {"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.float32 - ).untag("baz"), + "a": ( + pz.nx.zeros( + {"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.float32 + ).untag("baz") + ), }, pattern={ "a": pz.chk.ArraySpec(shape=(5,), named_shape={"foo": 3, "bar": 4}) @@ -116,9 +118,11 @@ def test_bad_named_shapes(self): with self.assertRaisesWithLiteralMatch(pz.chk.StructureMismatchError, err): pz.chk.check_structure( value={ - "a": pz.nx.zeros( - {"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.float32 - ).untag("baz"), + "a": ( + pz.nx.zeros( + {"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.float32 + ).untag("baz") + ), "b": pz.nx.zeros({"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.int32), "c": pz.nx.zeros( {"foo": 3, "bar": 4, "baz": 5}, dtype=jnp.float32 @@ -400,6 +404,30 @@ def test_named_unpack_substitution_conflict(self): }, ) + def test_named_axis_dimension_vars_support_nonstr_axis_names(self): + batch_axis = pz.nx.TmpPosAxisMarker() + match = pz.chk.check_structure( + value={ + "a": pz.chk.ArraySpec( + named_shape={batch_axis: 3, "features": 2}, dtype=jnp.float32 + ), + "b": pz.chk.ArraySpec( + named_shape={batch_axis: 3, "outputs": 4}, dtype=jnp.float32 + ), + }, + pattern={ + "a": pz.chk.ArraySpec( + named_shape={**pz.chk.var("batch"), "features": 2}, + dtype=jnp.float32, + ), + "b": pz.chk.ArraySpec( + named_shape={**pz.chk.var("batch"), "outputs": 4}, + dtype=jnp.float32, + ), + }, + ) + self.assertEqual(dict(match), {"batch": {batch_axis: 3}}) + def test_named_axis_inconsistent_shapes(self): err = textwrap.dedent("""\ Mismatch while checking structures: