@@ -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
827812lapack_inverse <- function (A , scope , hoist , dest = NULL , context = " solve" ) {
0 commit comments