Skip to content

Commit 43ca0fb

Browse files
authored
update update_gradient_JTCJ_sparse (#1225)
1 parent 295dbac commit 43ca0fb

1 file changed

Lines changed: 49 additions & 55 deletions

File tree

‎mujoco_warp/_src/solver.py‎

Lines changed: 49 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -2453,7 +2453,7 @@ def update_gradient_JTCJ_sparse(
24532453
nblocks_perblock: int,
24542454
dim_block: int,
24552455
# Out:
2456-
h_out: wp.array3d(dtype=float),
2456+
ctx_h_out: wp.array3d(dtype=float),
24572457
):
24582458
conid_start, elementid = wp.tid()
24592459

@@ -2468,20 +2468,37 @@ def update_gradient_JTCJ_sparse(
24682468

24692469
worldid = contact_worldid_in[conid]
24702470
if ctx_done_in[worldid]:
2471-
return
2471+
continue
24722472

24732473
condim = contact_dim_in[conid]
24742474

24752475
if condim == 1:
2476-
return
2476+
continue
24772477

24782478
# check contact status
24792479
if contact_dist_in[conid] - contact_includemargin_in[conid] >= 0.0:
2480-
return
2480+
continue
24812481

24822482
efcid0 = contact_efc_address_in[conid, 0]
24832483
if efc_state_in[worldid, efcid0] != types.ConstraintState.CONE:
2484-
return
2484+
continue
2485+
2486+
# All dims share the same sparsity pattern. Scan colind once to find
2487+
# the sparse positions of dof1id and dof2id. Skip if either is absent.
2488+
rownnz = efc_J_rownnz_in[worldid, efcid0]
2489+
rowadr0 = efc_J_rowadr_in[worldid, efcid0]
2490+
pos1 = int(-1)
2491+
pos2 = int(-1)
2492+
for k in range(rownnz):
2493+
col = efc_J_colind_in[worldid, 0, rowadr0 + k]
2494+
if col == dof1id:
2495+
pos1 = k
2496+
if col == dof2id:
2497+
pos2 = k
2498+
if pos1 >= 0 and pos2 >= 0:
2499+
break
2500+
if pos1 < 0 or pos2 < 0:
2501+
continue
24852502

24862503
fri = contact_friction_in[conid]
24872504
mu = fri[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]]
@@ -2490,7 +2507,7 @@ def update_gradient_JTCJ_sparse(
24902507
dm = math.safe_div(efc_D_in[worldid, efcid0], mu2 * (1.0 + mu2))
24912508

24922509
if dm == 0.0:
2493-
return
2510+
continue
24942511

24952512
n = ctx_Jaref_in[worldid, efcid0] * mu
24962513
u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0)
@@ -2509,89 +2526,66 @@ def update_gradient_JTCJ_sparse(
25092526
t = wp.max(t, types.MJ_MINVAL)
25102527
ttt = wp.max(t * t * t, types.MJ_MINVAL)
25112528

2529+
# Precompute common subexpressions.
2530+
mu_over_t = math.safe_div(mu, t)
2531+
mu_n_over_ttt = mu * math.safe_div(n, ttt)
2532+
mu2_minus_mu_n_over_t = mu2 - mu * math.safe_div(n, t)
2533+
25122534
h = float(0.0)
25132535

25142536
for dim1id in range(condim):
25152537
if dim1id == 0:
2516-
efcid1 = efcid0
2538+
rowadr1 = rowadr0
2539+
dm_fri1 = dm * mu
25172540
else:
25182541
efcid1 = contact_efc_address_in[conid, dim1id]
2542+
rowadr1 = efc_J_rowadr_in[worldid, efcid1]
2543+
dm_fri1 = dm * fri[dim1id - 1]
25192544

2520-
# TODO(team): improve performance for sparse code path
2521-
rownnz1 = efc_J_rownnz_in[worldid, efcid1]
2522-
rowadr1 = efc_J_rowadr_in[worldid, efcid1]
2523-
2524-
efc_J11 = float(0.0)
2525-
efc_J12 = float(0.0)
2526-
for i1 in range(rownnz1):
2527-
sparseid1 = rowadr1 + i1
2528-
colind1 = efc_J_colind_in[worldid, 0, sparseid1]
2529-
if dof1id == colind1:
2530-
efc_J11 = efc_J_in[worldid, 0, sparseid1]
2531-
if dof2id == colind1:
2532-
efc_J12 = efc_J_in[worldid, 0, sparseid1]
2533-
if efc_J11 != 0.0 and efc_J12 != 0.0:
2534-
break
2545+
# Direct J reads using cached sparse positions.
2546+
efc_J11 = efc_J_in[worldid, 0, rowadr1 + pos1]
2547+
efc_J12 = efc_J_in[worldid, 0, rowadr1 + pos2]
25352548

25362549
ui = u[dim1id]
25372550

25382551
for dim2id in range(0, dim1id + 1):
25392552
if dim2id == 0:
2540-
efcid2 = efcid0
2553+
rowadr2 = rowadr0
2554+
dm_fri12 = dm_fri1 * mu
25412555
else:
25422556
efcid2 = contact_efc_address_in[conid, dim2id]
2557+
rowadr2 = efc_J_rowadr_in[worldid, efcid2]
2558+
dm_fri12 = dm_fri1 * fri[dim2id - 1]
25432559

2544-
rownnz2 = efc_J_rownnz_in[worldid, efcid2]
2545-
rowadr2 = efc_J_rowadr_in[worldid, efcid2]
2546-
2547-
efc_J21 = float(0.0)
2548-
efc_J22 = float(0.0)
2549-
for i2 in range(rownnz2):
2550-
sparseid2 = rowadr2 + i2
2551-
colind2 = efc_J_colind_in[worldid, 0, sparseid2]
2552-
if dof1id == colind2:
2553-
efc_J21 = efc_J_in[worldid, 0, sparseid2]
2554-
if dof2id == colind2:
2555-
efc_J22 = efc_J_in[worldid, 0, sparseid2]
2556-
if efc_J21 != 0.0 and efc_J22 != 0.0:
2557-
break
2560+
# Direct J reads using cached sparse positions.
2561+
efc_J21 = efc_J_in[worldid, 0, rowadr2 + pos1]
2562+
efc_J22 = efc_J_in[worldid, 0, rowadr2 + pos2]
25582563

25592564
uj = u[dim2id]
25602565

25612566
# set first row/column: (1, -mu/t * u)
25622567
if dim1id == 0 and dim2id == 0:
25632568
hcone = 1.0
25642569
elif dim1id == 0:
2565-
hcone = -math.safe_div(mu, t) * uj
2570+
hcone = -mu_over_t * uj
25662571
elif dim2id == 0:
2567-
hcone = -math.safe_div(mu, t) * ui
2572+
hcone = -mu_over_t * ui
25682573
else:
2569-
hcone = mu * math.safe_div(n, ttt) * ui * uj
2574+
hcone = mu_n_over_ttt * ui * uj
25702575

25712576
# add to diagonal: mu^2 - mu * n / t
25722577
if dim1id == dim2id:
2573-
hcone += mu2 - mu * math.safe_div(n, t)
2574-
2575-
# pre and post multiply by diag(mu, friction) scale by dm
2576-
if dim1id == 0:
2577-
fri1 = mu
2578-
else:
2579-
fri1 = fri[dim1id - 1]
2578+
hcone += mu2_minus_mu_n_over_t
25802579

2581-
if dim2id == 0:
2582-
fri2 = mu
2583-
else:
2584-
fri2 = fri[dim2id - 1]
2585-
2586-
hcone *= dm * fri1 * fri2
2580+
hcone *= dm_fri12
25872581

25882582
if hcone != 0.0:
25892583
h += hcone * efc_J11 * efc_J22
25902584

25912585
if dim1id != dim2id:
25922586
h += hcone * efc_J12 * efc_J21
25932587

2594-
h_out[worldid, dof1id, dof2id] += h
2588+
ctx_h_out[worldid, dof1id, dof2id] += h
25952589

25962590

25972591
@wp.kernel
@@ -3003,7 +2997,7 @@ def _set_h_qM_dense(
30032997
if SPARSE_CONSTRAINT_JACOBIAN:
30042998
wp.launch(
30052999
update_gradient_JTCJ_sparse,
3006-
dim=(d.naconmax, m.dof_tri_row.size),
3000+
dim=(dim_block, m.dof_tri_row.size),
30073001
inputs=[
30083002
m.opt.impratio_invsqrt,
30093003
m.dof_tri_row,

0 commit comments

Comments
 (0)