Return the total non-zeros from mju_sqrMatTDSparseCount.
PiperOrigin-RevId: 772414653 Change-Id: Ic1cac868af0034196f33d08b2456be77c72cbd50
This commit is contained in:
committed by
Copybara-Service
parent
ea949f1cee
commit
998769a995
@@ -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));
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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];
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user