@@ -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