Skip to content
4 changes: 2 additions & 2 deletions pytato/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def set_debug_enabled(flag: bool) -> None:
Einsum, Stack, Concatenate, AxisPermutation,
IndexBase, Roll, IndexRemappingBase, BasicIndex,
AdvancedIndexInContiguousAxes, AdvancedIndexInNoncontiguousAxes,
SizeParam, Axis,
SizeParam, Axis, ReductionDescriptor,

make_dict_of_named_arrays,
make_placeholder, make_size_param, make_data_wrapper,
Expand Down Expand Up @@ -113,7 +113,7 @@ def set_debug_enabled(flag: bool) -> None:
"Stack", "Concatenate", "AxisPermutation",
"IndexBase", "Roll", "IndexRemappingBase",
"AdvancedIndexInContiguousAxes", "AdvancedIndexInNoncontiguousAxes",
"BasicIndex", "SizeParam", "Axis",
"BasicIndex", "SizeParam", "Axis", "ReductionDescriptor",

"make_dict_of_named_arrays", "make_placeholder", "make_size_param",
"make_data_wrapper", "einsum",
Expand Down
7 changes: 4 additions & 3 deletions pytato/analysis/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,7 +227,8 @@ def is_einsum_similar_to_subscript(expr: Einsum, subscripts: str) -> bool:
would compute the same result as *expr*.
"""

from pytato.array import ElementwiseAxis, ReductionAxis, EinsumAxisDescriptor
from pytato.array import (EinsumElementwiseAxis, EinsumReductionAxis,
EinsumAxisDescriptor)

if not isinstance(expr, Einsum):
raise TypeError(f"{expr} expected to be Einsum, got {type(expr)}.")
Expand All @@ -243,7 +244,7 @@ def is_einsum_similar_to_subscript(expr: Einsum, subscripts: str) -> bool:

for idim, idx in enumerate(_get_indices_from_input_subscript(out_spec,
is_output=True)):
index_to_descrs[idx] = ElementwiseAxis(idim)
index_to_descrs[idx] = EinsumElementwiseAxis(idim)

if len(in_spec.split(",")) != len(expr.args):
return False
Expand All @@ -263,7 +264,7 @@ def is_einsum_similar_to_subscript(expr: Einsum, subscripts: str) -> bool:
if index_to_descrs[idx] != access_descr:
return False
except KeyError:
if not isinstance(access_descr, ReductionAxis):
if not isinstance(access_descr, EinsumReductionAxis):
return False
index_to_descrs[idx] = access_descr

Expand Down
Loading