Skip to content

Commit 7c7c0eb

Browse files
committed
Unify lapack_solve's branch structure and output selection
The dgesv branch ended in return(), leaving the qr.solve condition that followed always-true (and the function textually able to fall off the end). The branches are a plain if/else now, converging on one tail. The identical output-target selection is one shared spelling (dest_usable), still evaluated at each branch's original write point so declaration order and error order in the emitted block are unchanged. Review finding (fable-final-review.md #2); no behavior change.
1 parent fbff0f0 commit 7c7c0eb

1 file changed

Lines changed: 31 additions & 46 deletions

File tree

R/r2f-matrix-blas.R

Lines changed: 31 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -634,6 +634,21 @@ lapack_solve <- function(
634634

635635
nrhs <- if (b_rank == 1L) 1L else dim_or_one(B, 2L)
636636

637+
# Both lowerings write a solution shaped by R's contract: length follows
638+
# ncol(a), width follows the right-hand side. Each branch resolves the
639+
# output target at its own write point (declaration order matters for
640+
# the emitted block) with the one shared spelling below.
641+
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
642+
dest_usable <- function() {
643+
can_use_output(
644+
dest,
645+
input_names = c(A_name, B_input_name),
646+
expected_dims = expected_dims,
647+
context = context,
648+
allow_alias = B_input_name
649+
)
650+
}
651+
637652
# R's solve() requires a square `a`; least squares is qr.solve()'s job.
638653
# Statically rectangular `a` is a compile error, symbolic dims get a
639654
# runtime guard before the dgesv call. (A rectangular `a` used to fall
@@ -643,24 +658,13 @@ lapack_solve <- function(
643658
A_work <- hoist$declare_tmp(mode = "double", dims = list(m, m))
644659
hoist$emit(glue("{A_work@name} = {A_name}"))
645660

646-
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
647-
writes_to_dest <- FALSE
648-
if (
649-
can_use_output(
650-
dest,
651-
input_names = c(A_name, B_input_name),
652-
expected_dims = expected_dims,
653-
context = context,
654-
allow_alias = B_input_name
655-
)
656-
) {
657-
out_var <- dest
658-
out_name <- dest@name
659-
writes_to_dest <- TRUE
661+
use_dest <- dest_usable()
662+
out_var <- if (use_dest) {
663+
dest
660664
} else {
661-
out_var <- hoist$declare_tmp(mode = "double", dims = expected_dims)
662-
out_name <- out_var@name
665+
hoist$declare_tmp(mode = "double", dims = expected_dims)
663666
}
667+
out_name <- out_var@name
664668
# The output length follows ncol(a) (R's contract) while `b` follows
665669
# nrow(a); the two are only runtime-equal. When ncol is statically 1
666670
# the output declares as a scalar, so a symbolic-length `b` must be
@@ -691,15 +695,7 @@ lapack_solve <- function(
691695
hoist = hoist,
692696
scope = scope
693697
)
694-
695-
out <- Fortran(out_name, out_var)
696-
if (writes_to_dest) {
697-
out@writes_to_dest <- TRUE
698-
}
699-
return(out)
700-
}
701-
702-
if (identical(context, "qr.solve")) {
698+
} else {
703699
A_work <- hoist$declare_tmp(mode = "double", dims = list(m, n))
704700
hoist$emit(glue("{A_work@name} = {A_name}"))
705701

@@ -762,24 +758,13 @@ end do"
762758
scope = scope
763759
)
764760

765-
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
766-
writes_to_dest <- FALSE
767-
if (
768-
can_use_output(
769-
dest,
770-
input_names = c(A_name, B_input_name),
771-
expected_dims = expected_dims,
772-
context = context,
773-
allow_alias = B_input_name
774-
)
775-
) {
776-
out_var <- dest
777-
out_name <- dest@name
778-
writes_to_dest <- TRUE
761+
use_dest <- dest_usable()
762+
out_var <- if (use_dest) {
763+
dest
779764
} else {
780-
out_var <- hoist$declare_tmp(mode = "double", dims = expected_dims)
781-
out_name <- out_var@name
765+
hoist$declare_tmp(mode = "double", dims = expected_dims)
782766
}
767+
out_name <- out_var@name
783768

784769
if (passes_as_scalar(out_var)) {
785770
hoist$emit(glue("{out_name} = {coef_work@name}(1, 1)"))
@@ -806,13 +791,13 @@ end do"
806791
))
807792
}
808793
}
794+
}
809795

810-
out <- Fortran(out_name, out_var)
811-
if (writes_to_dest) {
812-
out@writes_to_dest <- TRUE
813-
}
814-
return(out)
796+
out <- Fortran(out_name, out_var)
797+
if (use_dest) {
798+
out@writes_to_dest <- TRUE
815799
}
800+
out
816801
}
817802

818803
lapack_inverse <- function(A, scope, hoist, dest = NULL, context = "solve") {

0 commit comments

Comments
 (0)