-
Notifications
You must be signed in to change notification settings - Fork 8
Forward rules for QR/LQ #283
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
13 commits
Select commit
Hold shift + click to select a range
8fab5a5
Forward rules for QR/LQ
kshyatt 5aa1790
Calm Enzyme on 1.10
kshyatt e19beb5
Use Mooncake branch on CUDA for now
kshyatt 4281ef5
Fix lq
kshyatt 30e7344
Apply batched suggestions from code review
kshyatt 76221d6
Formatter
a613505
Fix typos from suggestions
847b5c1
Fix in case of aliases
kshyatt a5a0eeb
Update Project.tomls
d723c58
Remove extraneous inactive_type line
kshyatt 1b5b4a8
Remove unnecessary reinstantiation of A
kshyatt 44a2308
Add todo comment to lq
kshyatt 3f9e51b
Add a TODO comment to qr pushforward also
kshyatt File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Some comments aren't visible on the classic Files Changed page.
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,82 @@ | ||
| """ | ||
| lq_pushforward!( | ||
| ΔA, A, LQ, ΔLQ; | ||
| rank_atol::Real = default_pullback_rank_atol(LQ[1]) | ||
| ) | ||
|
|
||
| Compute the pushforward `ΔLQ` of the LQ decomposition `LQ` of `lq_compact(A; | ||
| positive = true)` or `lq_full(A; positive = true)` given the tangent `ΔA` of `A`. | ||
|
|
||
| If the original matrix `A` is rank-deficient (rank `r < min(size(A)...)`), only the first `r` | ||
| rows of `Q` and the first `r` columns of `L` are differentiable, and the tangents of the | ||
| remaining rows of `Q` and columns of `L` are set to zero. Similarly, for `lq_full` the extra | ||
| rows of `Q` are only determined up to a unitary rotation, and only their gauge-invariant | ||
| tangent component along the first `r` rows of `Q` is computed. | ||
|
|
||
| See also [`qr_pushforward!`](@ref). | ||
| """ | ||
| function lq_pushforward!( | ||
| ΔA, A, LQ, ΔLQ; | ||
| rank_atol::Real = default_pullback_rank_atol(LQ[1]), kwargs... | ||
| ) | ||
| L, Q = LQ | ||
| ΔL, ΔQ = ΔLQ | ||
| m = size(L, 1) | ||
| n = size(Q, 2) | ||
| minmn = min(m, n) | ||
| p = lq_rank(L; rank_atol) | ||
| (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of L*Q ($m, $n)")) | ||
|
|
||
| Q₁ = view(Q, 1:p, :) | ||
| # Julia 1.13 `ldiv!` checks `istriu` on the parent, which | ||
| # falls back to scalar indexing for a view of a GPU array | ||
| L₁₁ = LowerTriangular(L[1:p, 1:p]) | ||
| L₂₁ = view(L, (p + 1):m, 1:p) | ||
|
|
||
| ΔA₁ = view(ΔA, 1:p, :) | ||
| ΔA₂ = view(ΔA, (p + 1):m, :) | ||
|
|
||
| # compute everything from ΔA before writing to ΔL and ΔQ, which may alias it | ||
| # (e.g. `lq_compact!` of a `Diagonal` returns `Q === A`) | ||
| # TODO rework this into two versions, one for `Q === A` and one for | ||
| # `Q !== A`, to minimize allocations | ||
| ΔQ₁ = L₁₁ \ ΔA₁ | ||
| ΔQ₁Q₁ᴴ = ΔQ₁ * Q₁' | ||
| M = ΔQ₁Q₁ᴴ + ΔQ₁Q₁ᴴ' | ||
| diagview(M) ./= 2 | ||
| view(M, uppertriangularind(M)) .= zero(eltype(M)) | ||
| ΔL₁₁ = L₁₁ * M | ||
| ΔQ₁ = mul!(ΔQ₁, M, Q₁, -1, 1) | ||
| ΔL₂₁ = ΔA₂ * Q₁' | ||
| ΔL₂₁ = mul!(ΔL₂₁, L₂₁, ΔQ₁ * Q₁', -1, 1) | ||
|
|
||
| zero!(ΔL) | ||
| zero!(ΔQ) | ||
| view(ΔQ, 1:p, :) .= ΔQ₁ | ||
| view(ΔL, 1:p, 1:p) .= ΔL₁₁ | ||
| view(ΔL, (p + 1):m, 1:p) .= ΔL₂₁ | ||
| if p == minmn && size(Q, 1) > minmn | ||
| Q₃ = view(Q, (minmn + 1):size(Q, 1), :) | ||
| ΔQ₃ = view(ΔQ, (minmn + 1):size(Q, 1), :) | ||
| mul!(ΔQ₃, Q₃ * ΔQ₁', Q₁, -1, 0) | ||
| end | ||
| return ΔL, ΔQ | ||
| end | ||
|
|
||
| """ | ||
| lq_null_pushforward!(ΔA, A, Nᴴ, ΔNᴴ; kwargs...) | ||
|
|
||
| Compute the pushforward `ΔNᴴ` of the left nullspace basis `Nᴴ` of `lq_null(A)` | ||
| given the tangent `ΔA` of `A`. | ||
|
|
||
| See also [`lq_pushforward!`](@ref). | ||
| """ | ||
| function lq_null_pushforward!(ΔA, A, Nᴴ, ΔNᴴ; kwargs...) | ||
| if size(Nᴴ, 1) == 0 | ||
| zero!(ΔNᴴ) | ||
| return ΔNᴴ | ||
| end | ||
| L, Q = lq_compact(A; positive = true) | ||
| ΔQNᴴ = ldiv!(LowerTriangular(L), ΔA * Nᴴ') | ||
| return mul!(ΔNᴴ, ΔQNᴴ', Q, -1, 0) | ||
| end |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,77 @@ | ||
| """ | ||
| qr_pushforward!( | ||
| ΔA, A, QR, ΔQR; | ||
| rank_atol::Real = default_pullback_rank_atol(QR[2]) | ||
| ) | ||
|
|
||
| Computes the pushforward `ΔQR` of the QR decomposition `QR` of `qr_compact(A; | ||
| positive = true)` or `qr_full(A; positive = true)` given the tangent `ΔA` of `A`. | ||
|
|
||
| If the original matrix `A` is rank-deficient (rank `r < min(size(A)...)`), only the first `r` | ||
| columns of `Q` and the first `r` rows of `R` are differentiable, and the tangents of the | ||
| remaining columns of `Q` and rows of `R` are set to zero. Similarly, for `qr_full` the extra | ||
| columns of `Q` are only determined up to a unitary rotation, and only their gauge-invariant | ||
| tangent component along the first `r` columns of `Q` is computed. | ||
| """ | ||
| function qr_pushforward!( | ||
| ΔA, A, QR, ΔQR; | ||
| rank_atol::Real = default_pullback_rank_atol(QR[2]), kwargs... | ||
| ) | ||
| Q, R = QR | ||
| ΔQ, ΔR = ΔQR | ||
| m = size(Q, 1) | ||
| n = size(R, 2) | ||
| minmn = min(m, n) | ||
| p = qr_rank(R; rank_atol) | ||
| (m, n) == size(ΔA) || throw(DimensionMismatch("size of ΔA ($(size(ΔA))) does not match size of Q*R ($m, $n)")) | ||
|
|
||
| # TODO rework this into two versions, one for `Q === A` and one for | ||
| # `Q !== A`, to minimize allocations | ||
| Q₁ = view(Q, :, 1:p) | ||
| R₁₁ = UpperTriangular(view(R, 1:p, 1:p)) | ||
| R₁₂ = view(R, 1:p, (p + 1):n) | ||
|
|
||
| ΔA₁ = view(ΔA, :, 1:p) | ||
| ΔA₂ = view(ΔA, :, (p + 1):n) | ||
|
|
||
| ΔQ₁ = ΔA₁ / R₁₁ | ||
| Q₁ᴴΔQ₁ = Q₁' * ΔQ₁ | ||
| M = Q₁ᴴΔQ₁ + Q₁ᴴΔQ₁' | ||
| diagview(M) ./= 2 | ||
| view(M, lowertriangularind(M)) .= zero(eltype(M)) | ||
| ΔR₁₁ = M * R₁₁ | ||
| ΔQ₁ = mul!(ΔQ₁, Q₁, M, -1, 1) | ||
| ΔR₁₂ = Q₁' * ΔA₂ | ||
| ΔR₁₂ = mul!(ΔR₁₂, Q₁' * ΔQ₁, R₁₂, -1, 1) | ||
|
|
||
| zero!(ΔQ) | ||
| zero!(ΔR) | ||
| view(ΔQ, :, 1:p) .= ΔQ₁ | ||
| view(ΔR, 1:p, 1:p) .= ΔR₁₁ | ||
| view(ΔR, 1:p, (p + 1):n) .= ΔR₁₂ | ||
| if p == minmn && size(Q, 2) > minmn | ||
| # extra columns in the case of qr_full, orthogonality to Q₁ fixes their component along Q₁ | ||
| Q₃ = view(Q, :, (minmn + 1):size(Q, 2)) | ||
| ΔQ₃ = view(ΔQ, :, (minmn + 1):size(Q, 2)) | ||
| mul!(ΔQ₃, Q₁, ΔQ₁' * Q₃, -1, 0) | ||
| end | ||
| return ΔQ, ΔR | ||
| end | ||
|
|
||
| """ | ||
| qr_null_pushforward!(ΔA, A, N, ΔN; kwargs...) | ||
|
|
||
| Compute the pushforward `ΔN` of the nullspace basis `N` of `qr_null(A)` given the | ||
| tangent `ΔA` of `A`. | ||
|
|
||
| See also [`qr_pushforward!`](@ref). | ||
| """ | ||
| function qr_null_pushforward!(ΔA, A, N, ΔN; kwargs...) | ||
| if size(N, 2) == 0 | ||
| zero!(ΔN) | ||
| return ΔN | ||
| end | ||
| Q, R = qr_compact(A; positive = true) | ||
| NᴴΔQ = rdiv!(N' * ΔA, UpperTriangular(R)) | ||
| return mul!(ΔN, Q, NᴴΔQ', -1, 0) | ||
| end | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.