Вычисление градиента с использованием JAX функции, которая выводит список массивовPython

Программы на Python
Anonymous
Вычисление градиента с использованием JAX функции, которая выводит список массивов

Сообщение Anonymous »

У меня есть функция, которая возвращает список массивов, и мне нужно найти ее производную по одному параметру. Например, предположим, что у нас есть

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

def fun(x):
...
return [a,b,c]
где a,b,c и d — многомерные массивы (например, реальные массивы 2 на 2 на 2). Теперь я хочу получить [da/dx, db/dx, dc/dx]. Под db/dx я имею в виду, что хочу получить производную каждого элемента массива a:222 по x, поэтому все da/dx, db/dx, dc/dx равны 2< em>22 массива.
Я впервые использую дифференцирование JAX, и большинство примеров, которые я нахожу в Интернете, посвящены функциям со скалярным выводом.
Из моего поиска я понял, что один из способов найти это - это получить градиент каждого скаляра во всех этих массивах по одному (вероятно, это можно сделать быстрее, используя vmap). Есть ли другой способ, более быстрый? Я думаю, что JAX.jacobian может помочь, но мне трудно найти его документацию, чтобы увидеть, что именно делает эта функция. Мы очень ценим любую помощь.
Теперь я попробовал JAX.jacobian на простых примерах, и он дал мне ожидаемый ответ. Это меня немного успокаивает, но я хотел бы найти официальную документацию или подтверждение от других, что это правильный способ сделать это, и он делает то, что я от него ожидаю.

Подробнее здесь: https://stackoverflow.com/questions/790 ... -of-arrays

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