Skip to content

Commit 52bf76a

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 4620eaf commit 52bf76a

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
@@ -643,6 +643,21 @@ lapack_solve <- function(
643643

644644
nrhs <- if (b_rank == 1L) 1L else dim_or_one(B, 2L)
645645

646+
# Both lowerings write a solution shaped by R's contract: length follows
647+
# ncol(a), width follows the right-hand side. Each branch resolves the
648+
# output target at its own write point (declaration order matters for
649+
# the emitted block) with the one shared spelling below.
650+
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
651+
dest_usable <- function() {
652+
can_use_output(
653+
dest,
654+
input_names = c(A_name, B_input_name),
655+
expected_dims = expected_dims,
656+
context = context,
657+
allow_alias = B_input_name
658+
)
659+
}
660+
646661
# R's solve() requires a square `a`; least squares is qr.solve()'s job.
647662
# Statically rectangular `a` is a compile error, symbolic dims get a
648663
# runtime guard before the dgesv call. (A rectangular `a` used to fall
@@ -652,24 +667,13 @@ lapack_solve <- function(
652667
A_work <- hoist$declare_tmp(mode = "double", dims = list(m, m))
653668
hoist$emit(glue("{A_work@name} = {A_name}"))
654669

655-
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
656-
writes_to_dest <- FALSE
657-
if (
658-
can_use_output(
659-
dest,
660-
input_names = c(A_name, B_input_name),
661-
expected_dims = expected_dims,
662-
context = context,
663-
allow_alias = B_input_name
664-
)
665-
) {
666-
out_var <- dest
667-
out_name <- dest@name
668-
writes_to_dest <- TRUE
670+
use_dest <- dest_usable()
671+
out_var <- if (use_dest) {
672+
dest
669673
} else {
670-
out_var <- hoist$declare_tmp(mode = "double", dims = expected_dims)
671-
out_name <- out_var@name
674+
hoist$declare_tmp(mode = "double", dims = expected_dims)
672675
}
676+
out_name <- out_var@name
673677
# The output length follows ncol(a) (R's contract) while `b` follows
674678
# nrow(a); the two are only runtime-equal. When ncol is statically 1
675679
# the output declares as a scalar, so a symbolic-length `b` must be
@@ -700,15 +704,7 @@ lapack_solve <- function(
700704
hoist = hoist,
701705
scope = scope
702706
)
703-
704-
out <- Fortran(out_name, out_var)
705-
if (writes_to_dest) {
706-
out@writes_to_dest <- TRUE
707-
}
708-
return(out)
709-
}
710-
711-
if (identical(context, "qr.solve")) {
707+
} else {
712708
A_work <- hoist$declare_tmp(mode = "double", dims = list(m, n))
713709
hoist$emit(glue("{A_work@name} = {A_name}"))
714710

@@ -771,24 +767,13 @@ end do"
771767
scope = scope
772768
)
773769

774-
expected_dims <- if (b_rank == 1L) list(n) else list(n, nrhs)
775-
writes_to_dest <- FALSE
776-
if (
777-
can_use_output(
778-
dest,
779-
input_names = c(A_name, B_input_name),
780-
expected_dims = expected_dims,
781-
context = context,
782-
allow_alias = B_input_name
783-
)
784-
) {
785-
out_var <- dest
786-
out_name <- dest@name
787-
writes_to_dest <- TRUE
770+
use_dest <- dest_usable()
771+
out_var <- if (use_dest) {
772+
dest
788773
} else {
789-
out_var <- hoist$declare_tmp(mode = "double", dims = expected_dims)
790-
out_name <- out_var@name
774+
hoist$declare_tmp(mode = "double", dims = expected_dims)
791775
}
776+
out_name <- out_var@name
792777

793778
if (passes_as_scalar(out_var)) {
794779
hoist$emit(glue("{out_name} = {coef_work@name}(1, 1)"))
@@ -815,13 +800,13 @@ end do"
815800
))
816801
}
817802
}
803+
}
818804

819-
out <- Fortran(out_name, out_var)
820-
if (writes_to_dest) {
821-
out@writes_to_dest <- TRUE
822-
}
823-
return(out)
805+
out <- Fortran(out_name, out_var)
806+
if (use_dest) {
807+
out@writes_to_dest <- TRUE
824808
}
809+
out
825810
}
826811

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

0 commit comments

Comments
 (0)