torch.func
torch.func, ранее известный как «functorch», представляет собой аналогичные JAX композируемые преобразования функций для PyTorch.
Примечание
Эта библиотека в настоящее время находится в бета-версии. Это означает, что функции в целом работают (если не указано иное), и мы (команда PyTorch) стремимся к развитию этой библиотеки. Однако API могут быть изменены на основании обратной связи пользователей, и у нас нет полной поддержки всех операций PyTorch.
Если у вас есть предложения по API или сценариям использования, которые вы хотели бы видеть, откройте вопрос на GitHub или свяжитесь с нами. Мы с удовольствием узнаем о том, как вы используете эту библиотеку.
Что такое композируемые преобразования функций?
- «Преобразование функции» — это функция высшего порядка, которая принимает численную функцию и возвращает новую функцию, вычисляющую другую величину.
-
torch.funcимеет преобразования автоматического дифференцирования (grad(f)возвращает функцию, вычисляющую градиентf), преобразование векторизации/пакетирования (vmap(f)возвращает функцию, вычисляющуюfпо пакетам входных данных) и другие. - Эти преобразования функций могут комбинироваться друг с другом произвольно. Например, комбинирование
vmap(grad(f))вычисляет величину, называемую градиентами по каждому образцу, которую PyTorch сегодня не может эффективно вычислить.
Зачем нужны композируемые преобразования функций?
Существует ряд сценариев использования, которые трудно реализовать в PyTorch сегодня:
- вычисление градиентов по каждому образцу (или других величин по каждому образцу)
- запуск ансамблей моделей на одном компьютере
- эффективное объединение задач в цикле внутреннего уровня MAML
- эффективное вычисление Якобианов и Гессианов
- эффективное вычисление пакетированных Якобианов и Гессианов
Комбинирование vmap(), grad() и vjp() преобразований позволяет нам выразить вышеперечисленное, не разрабатывая отдельную подсистему для каждого случая. Эта идея композируемых преобразований функций заимствована из фреймворка JAX.
Подробнее
- Быстрый обзор torch.func
- Справочник API torch.func
- Ограничения в пользовательском интерфейсе
- Миграция с functorch на torch.func
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/func.html