Add additional type check in len wrapper hack to prevent issues with Tensors in list

This commit is contained in:
Jedrzej Kosinski
2024-08-04 21:08:43 -05:00
parent c4aa684cb3
commit 17b70c421a
+1 -1
View File
@@ -39,7 +39,7 @@ def wrapper_len_factory(orig_len: Callable) -> Callable:
def wrapper_len(*args, **kwargs):
cond_or_uncond = args[0]
real_length = orig_len(*args, **kwargs)
if real_length > 0 and type(cond_or_uncond) == list and (cond_or_uncond[0] in [0, 1]):
if real_length > 0 and type(cond_or_uncond) == list and isinstance(cond_or_uncond[0], int) and (cond_or_uncond[0] in [0, 1]):
try:
to_return = IntWithCondOrUncond(real_length)
setattr(to_return, "cond_or_uncond", cond_or_uncond)