Use ChainRules for operators
まだ誰も着手していません。
評価
- 難易度
- 5/5
- 見積もり時間
- 1週間以上
- 初心者へのやさしさ
- 25/100
- issue の種類
- リファクタリング
- 明瞭さ
- 説明が足りない
- 活発さ
- 停滞
- 技術スタック
- julia
調査の方向性
まず src/Nonlinear/univariate_expressions_generator.jl と、src/Nonlinear/operators.jl:570-582 周辺の operator dispatch を読みます。ReverseAD が classical、registered、user-defined の各 operator をどのように表現しているかを追跡し、その流れを提案されている ChainRules-style dispatch と fallback と比較します。完了の条件は、設計が決定され、operator handling が更新され、サポート対象のケースで現在懸念されている type instability を回避できることを示す証拠があることです。
索引モデルが issue の本文から書いたものです。
説明
Currently, the approach in ReverseAD is to generate the symbolic expression of the first and second-order derivatives for classical univariate functions using Calculus
https://github.com/jump-dev/MathOptInterface.jl/blob/100eab2e669e73689e1dc214391d97c24402e35c/src/Nonlinear/univariate_expressions_generator.jl
Then, given a representation of the operator as an Int, we do an hard-coded binary search to evaluate a O(log(n)) number of Int comparison instead of a O(n) number of comparison:
https://github.com/jump-dev/MathOptInterface.jl/blob/100eab2e669e73689e1dc214391d97c24402e35c/src/Nonlinear/operators.jl#L570-L582
I'm wondering whether we could get closer to ChainRules instead like other Julia AD framework.
The naive way to do this would be
op = :tanh
f = eval(op)
value_and_derivative(f, 1)
The issue is that, because the value of op is discovered at run-time, the type of f is type-unstable.
But we can use the same trick with the if-else and do
if op == :tanh
value_and_derivative(tanh, x)
elseif op == :tan
value_and_derivative(tan, x)
elseif ...
else
value_and_derivative(eval(op), x)
end
Again, we can do a binary search instead of just a list of if-else.
So, for a fixed number of symbols, we avoid the type-instability thanks to the if-else and we have a fallback for the other ones with the eval.
That would also mean that for registered functions, we need to implement a method and just rely on multiple dispatch instead of adding an operators to the list of user-defined operators, user-defined operators already trigger a type-instability when they are called anyway.
- 主要言語
- Julia
- スター
- 6
- フォーク
- 0
- 平均マージ
- 6日 22時間
- マージ済み PR(30日)
- 4
環境構築
このプロジェクトには開発コンテナ、Dockerfile、コントリビューションガイドがありません。まず README を読み、一般的な手順ははじめてのコントリビューションガイドを参照してください。
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- issue 番号を参照したプルリクエストを送ります。
blegat/ArrayDiff.jl のほかの issue
-
難易度 2/5 1〜3時間 初心者へのやさしさ 76/100
blegat/ArrayDiff.jl#83 ·
-
Building matrixオープン
難易度 3/5 1〜2日 初心者へのやさしさ 45/100
blegat/ArrayDiff.jl#13 ·
blegat/ArrayDiff.jl の issue をすべて見る
似ている issue
-
難易度 2/5 1〜3時間 初心者へのやさしさ 75/100
-
難易度 2/5 1〜3時間 初心者へのやさしさ 84/100
メンテナーはふだん 1 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 92/100
JuliaSymbolics/SymbolicUtils.jl#1131 · コメント 1 件 ·
メンテナーはふだん 1 日以内に返信
-
documentation
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
メンテナーはふだん 2 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100