Error message for bad symbolic shapes is now useful.

This commit is contained in:
Patrick Kidger
2023-10-21 11:38:04 -07:00
parent 7f85a12a85
commit 2ab8286c81
+9 -6
View File
@@ -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, "<string>", "eval")
elem = _SymbolicDim(elem, broadcastable)
elem = _SymbolicDim(elem, broadcastable, elem_string)
dims.append(elem)
dims = tuple(dims)