Код выполняет условные проверки значений left и вправо как часть процесса обхода дерева - например. он будет пытаться пройти только по левой или правой ветви поддерева узла, если оно действительно существует. Выполнение проверок типа if (current_node.left не равно None) (или это должно быть jax.numpy.ological_not(current_node.left равно None) в Jax - я пробовал и то, и другое?) было это нормально, но после преобразования операторов if в jax.lax.cond(...) я получаю ошибку AttributeError: у объекта 'NoneType' нет атрибута 'left'.
Я думаю, ситуация может быть такой, как в следующем минимальном рабочем примере:
Код: Выделить всё
import jax
import jax.numpy as jnp
def my_func(val):
return 2*val
@jax.jit
def test_fn(a):
return jax.lax.cond(a is not None,
lambda: my_func(a),
lambda: 0)
print(test_fn(2)) # Prints 4
# in test_fn(), a has type
print(test_fn(None)) # TypeError: unsupported operand type(s) for *: 'int' and 'NoneType'
# in test_fn(), a has type
Правильно ли я понимаю и что мне делать с этой проблемой? Как ни странно, в минимальном рабочем примере кода также возникает проблема, когда декоратор @jax.jit опущен, что позволяет предположить, что обе ветви все еще отслеживаются.
Кстати, является ли древовидная структура «встроенной» в код Jax/XLA? Я заметил, что при использовании больших деревьев Jit-компиляция кода занимает больше времени, и это заставляет меня беспокоиться, что это может быть недопустимым подходом с очень большим количеством точек, которые мне нужно представить (около 14 000 000). Я бы использовал обычную реализацию kd-tree Scipy, но, к сожалению, она несовместима с Jax, и этого требует остальная часть моего кода. Я мог бы задать это как отдельный вопрос для ясности.
Подробнее здесь: https://stackoverflow.com/questions/787 ... ementation