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