Идея состоит в том, что этот векторизованный алгоритм будет проверять определенные критерии по 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.Возможно, я совершенно не могу работать с такими структурами данных через Jax, поскольку ошибка предполагает, что поддерживаются только числовые типы. Могу ли я каким-либо образом использовать pytrees, чтобы смягчить это ограничение?
Будет заманчиво предложить мне использовать простой список из jnp, но я использую Linkedlist только в качестве примера. простой(ой) структуры данных. Как упоминалось ранее, на самом деле я работаю над DCEL.
PS: код Linkedlist был взят с сайта GeeksForGeeks, так как я хотел быстро придумать минимальный рабочий пример.
Подробнее здесь: https://stackoverflow.com/questions/786 ... p-function