diff --git a/src/engine/engine_core_constraint.c b/src/engine/engine_core_constraint.c index 3e625f49..5b688bf4 100644 --- a/src/engine/engine_core_constraint.c +++ b/src/engine/engine_core_constraint.c @@ -2193,12 +2193,9 @@ void mj_projectConstraint(const mjModel* m, mjData* d) { } // pre-count A nonzeros (compute AR_rownnz, AR_rowadr) - mju_sqrMatTDSparseCount(d->efc_AR_rownnz, d->efc_AR_rowadr, nefc, - BT_rownnz, BT_rowadr, BT_colind, - B_rownnz, B_rowadr, B_colind, B_rowsuper, d, /*flg_upper=*/1); - - // nA = total number of nonzeros in A - d->nA = d->efc_AR_rownnz[nefc - 1] + d->efc_AR_rowadr[nefc - 1]; + d->nA = mju_sqrMatTDSparseCount(d->efc_AR_rownnz, d->efc_AR_rowadr, nefc, + BT_rownnz, BT_rowadr, BT_colind, + B_rownnz, B_rowadr, B_colind, B_rowsuper, d, /*flg_upper=*/1); // allocate A values and column indices on arena d->efc_AR = mj_arenaAllocByte(d, sizeof(mjtNum) * d->nA, _Alignof(mjtNum)); diff --git a/src/engine/engine_solver.c b/src/engine/engine_solver.c index 635fa6af..16286d4a 100644 --- a/src/engine/engine_solver.c +++ b/src/engine/engine_solver.c @@ -1533,15 +1533,14 @@ static void MakeHessian(mjData* d, mjCGContext* ctx) { // sparse if (ctx->is_sparse) { - // initialize Hessian rowadr, rownnz - mju_sqrMatTDSparseCount(ctx->H_rownnz, ctx->H_rowadr, nv, - ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind, - ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, - ctx->JT_rowsuper, d, /*flg_upper=*/0); + // initialize Hessian rowadr, rownnz; get total nonzeros + ctx->nH = mju_sqrMatTDSparseCount(ctx->H_rownnz, ctx->H_rowadr, nv, + ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind, + ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, + ctx->JT_rowsuper, d, /*flg_upper=*/0); - // add nC to Hessian total nonzeros (unavoidable overcounting since H_colind is still unknown) - ctx->nH = ctx->M_rowadr[nv - 1] + ctx->M_rownnz[nv - 1] + - ctx->H_rowadr[nv - 1] + ctx->H_rownnz[nv - 1]; + // add M nonzeros to Hessian total (unavoidable overcounting since H_colind is still unknown) + ctx->nH += ctx->M_rowadr[nv - 1] + ctx->M_rownnz[nv - 1]; // shift H row addresses to make room for C int shift = 0; diff --git a/src/engine/engine_util_sparse.c b/src/engine/engine_util_sparse.c index c9cf3bab..ad38edb7 100644 --- a/src/engine/engine_util_sparse.c +++ b/src/engine/engine_util_sparse.c @@ -636,11 +636,11 @@ void mju_superSparse(int nr, int* rowsuper, } -// precount res_rownnz and precompute res_rowadr for mju_sqrMatTDSparse -void mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, - const int* rownnz, const int* rowadr, const int* colind, - const int* rownnzT, const int* rowadrT, const int* colindT, - const int* rowsuperT, mjData* d, int flg_upper) { +// precount res_rownnz and precompute res_rowadr for mju_sqrMatTDSparse, return total non-zeros +int mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, + const int* rownnz, const int* rowadr, const int* colind, + const int* rownnzT, const int* rowadrT, const int* colindT, + const int* rowsuperT, mjData* d, int flg_upper) { mj_markStack(d); int* chain = mjSTACKALLOC(d, 2*nr, int); int nchain = 0; @@ -726,6 +726,8 @@ void mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, } mj_freeStack(d); + + return res_rowadr[nr-1] + res_rownnz[nr-1]; } diff --git a/src/engine/engine_util_sparse.h b/src/engine/engine_util_sparse.h index f6c2f7c2..4bf82472 100644 --- a/src/engine/engine_util_sparse.h +++ b/src/engine/engine_util_sparse.h @@ -120,11 +120,11 @@ MJAPI void mju_sqrMatTDSparse_row(mjtNum* res, const mjtNum* mat, const mjtNum* const int* colindT, const int* rowsuperT, mjData* d, int* diagind); -// precount res_rownnz and precompute res_rowadr for mju_sqrMatTDSparse -MJAPI void mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, - const int* rownnz, const int* rowadr, const int* colind, - const int* rownnzT, const int* rowadrT, const int* colindT, - const int* rowsuperT, mjData* d, int flg_upper); +// precount res_rownnz and precompute res_rowadr for mju_sqrMatTDSparse, return total non-zeros +MJAPI int mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, + const int* rownnz, const int* rowadr, const int* colind, + const int* rownnzT, const int* rowadrT, const int* colindT, + const int* rowsuperT, mjData* d, int flg_upper); // precompute res_rowadr for mju_sqrMatTDSparse using uncompressed memory MJAPI void mju_sqrMatTDUncompressedInit(int* res_rowadr, int nc);