Реализация векторизованной функции над LinkedLists с использованием функции Jax vmapPython

Программы на Python
Anonymous
Реализация векторизованной функции над LinkedLists с использованием функции Jax vmap

Сообщение Anonymous »

Пытаюсь реализовать векторизованную версию алгоритма (из вычислительной геометрии) с помощью Jax. Я создал минимальный рабочий пример с использованием LinkedList, чтобы конкретно выразить мой запрос (в противном случае я использую DCEL).
Идея состоит в том, что этот векторизованный алгоритм будет проверять определенные критерии по DCEL. . Для простоты я заменил эту «процедуру проверки критериев» на простой алгоритм суммирования.

Код: Выделить всё

import jax
from jax import vmap
import jax.numpy as jnp

class Node:

# Constructor to initialize the node object
def __init__(self, data):
self.data = data
self.next = None

class LinkedList:

def __init__(self):
self.head = None

def push(self, new_data):
new_node = Node(new_data)
new_node.next = self.head
self.head = new_node

def printList(self):
temp = self.head
while(temp):
print (temp.data,end=" ")
temp = temp.next

def summate(list) :
prev = None
current = list.head
sum = 0
while(current is not None):
sum += current.data
next = current.next
current = next
return sum

list1 = LinkedList()
list1.push(20)
list1.push(4)
list1.push(15)
list1.push(85)

list2 = LinkedList()
list2.push(19)
list2.push(13)
list2.push(2)
list2.push(13)

#list(map(summate, ([list1, list2])))

vmap(summate)(jnp.array([list1, list2]))

Я получаю следующую ошибку.

Код: Выделить всё

 TypeError: Value '' with dtype object is not a valid JAX array type. Only arrays of numeric types are supported by JAX.
Цель состоит в том, что если у меня есть набор, скажем, 10 000 Linkedlists, я смогу применить эту функцию суммирования к каждому LinkedList в векторизованном виде. Я реализовал то, что хотел, на базовом Python, но хочу сделать это на Jax, поскольку существует более крупная вероятностная функция, для которой я буду использовать эту подпроцедуру (это цепь Маркова).
Возможно, я совершенно не могу работать с такими структурами данных через Jax, поскольку ошибка предполагает, что поддерживаются только числовые типы. Могу ли я каким-либо образом использовать pytrees, чтобы смягчить это ограничение?
Будет заманчиво предложить мне использовать простой список из jnp, но я использую Linkedlist только в качестве примера. простой(ой) структуры данных. Как упоминалось ранее, на самом деле я работаю над DCEL.
PS: код Linkedlist был взят с сайта GeeksForGeeks, так как я хотел быстро придумать минимальный рабочий пример.

Подробнее здесь: https://stackoverflow.com/questions/786 ... p-function

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