Выполнение условных ветвей, вызывающих ошибки в Jax (реализация kd-tree)Python

Программы на Python
Anonymous
Выполнение условных ветвей, вызывающих ошибки в Jax (реализация kd-tree)

Сообщение Anonymous »

Я пишу kd-дерево в Jax и использую специально написанные объекты Node для элементов дерева. Каждый узел очень прост и имеет одно поле данных (для хранения числовых значений), а также левое и правое поля, которые являются ссылками на другие узлы. Листовой узел идентифицируется как узел, для которого поля left и right имеют значение None.
Код выполняет условные проверки значений 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 cond был обычным оператором if, my_func() не вызывался бы даже при < /code> имеет значение None, и ошибка не возникнет. Насколько я понимаю, Jax пытается отследить функцию, что означает, что все ветки выполняются, и это приводит к вызову my_func() с None (когда a имеет значение None), вызывая ошибку . Я считаю, что аналогичная ситуация возникает в моем древовидном коде, где условные ветви выполняются, даже если .left и/или .right равны None, а традиционный оператор if не будет привести к выполнению ветвей кода.
Правильно ли я понимаю и что мне делать с этой проблемой? Как ни странно, в минимальном рабочем примере кода также возникает проблема, когда декоратор @jax.jit опущен, что позволяет предположить, что обе ветви все еще отслеживаются.

Кстати, является ли древовидная структура «встроенной» в код Jax/XLA? Я заметил, что при использовании больших деревьев Jit-компиляция кода занимает больше времени, и это заставляет меня беспокоиться, что это может быть недопустимым подходом с очень большим количеством точек, которые мне нужно представить (около 14 000 000). Я бы использовал обычную реализацию kd-tree Scipy, но, к сожалению, она несовместима с Jax, и этого требует остальная часть моего кода. Я мог бы задать это как отдельный вопрос для ясности.

Подробнее здесь: https://stackoverflow.com/questions/787 ... ementation

Вернуться в «Python»