From 2ab8286c81b0f9c7340f40796cd78f1aa67b94b1 Mon Sep 17 00:00:00 2001 From: Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com> Date: Wed, 18 Oct 2023 16:58:56 -0700 Subject: [PATCH] Error message for bad symbolic shapes is now useful. --- jaxtyping/_array_types.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/jaxtyping/_array_types.py b/jaxtyping/_array_types.py index c31e74a..babdb76 100644 --- a/jaxtyping/_array_types.py +++ b/jaxtyping/_array_types.py @@ -86,9 +86,10 @@ class _FixedDim: class _SymbolicDim: - def __init__(self, expr, broadcastable): + def __init__(self, expr, broadcastable, elem_string): self.expr = expr self.broadcastable = broadcastable + self.elem_string = elem_string _AbstractDimOrVariadicDim = Union[ @@ -122,10 +123,11 @@ def _check_dims( eval_size = eval(cls_dim.expr, single_memo.copy()) except NameError as e: raise NameError( - f"Cannot process symbolic dimension '{cls_dim.expr}' as some " - "dimension names have not been processed. In practice you should " - "usually only use symbolic dimensions in annotations for return " - "types, referring only to dimensions annotated for arguments." + f"Cannot process symbolic dimension '{cls_dim.elem_string}' as " + "some dimension names have not been processed. In practice you " + "should usually only use symbolic dimensions in annotations for " + "return types, referring only to dimensions annotated for " + "arguments." ) from e if eval_size != obj_size: return False @@ -444,8 +446,9 @@ def _make_array(array_type, dim_str, dtypes, name): "Cannot have symbolic multiple-dimensions, e.g. " "`*foo+bar` is not allowed" ) + elem_string = elem elem = compile(elem, "", "eval") - elem = _SymbolicDim(elem, broadcastable) + elem = _SymbolicDim(elem, broadcastable, elem_string) dims.append(elem) dims = tuple(dims)