Skip to content

Commit 81c44ce

Browse files
committed
Add flag to disable snake benchmark external loads
1 parent 03375d3 commit 81c44ce

3 files changed

Lines changed: 81 additions & 58 deletions

File tree

benchmark/_jax_snake_common.py

Lines changed: 65 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -114,6 +114,7 @@ def build_cpu_sim(
114114
poisson_ratio: float,
115115
gravitational_acc: float,
116116
time_step: float,
117+
include_external_loads: bool = True,
117118
) -> tuple[MultiSnakeReferenceSimulator, list[ea.CosseratRod]]:
118119
b_coeff = default_b_coeff()
119120
normal = np.array([0.0, 1.0, 0.0], dtype=np.float64)
@@ -125,11 +126,12 @@ def build_cpu_sim(
125126

126127
sim = MultiSnakeReferenceSimulator()
127128
rods: list[ea.CosseratRod] = []
128-
ground_plane = ea.Plane(
129-
plane_origin=np.array([0.0, -base_length * 0.011, 0.0], dtype=np.float64),
130-
plane_normal=normal,
131-
)
132-
sim.append(ground_plane)
129+
if include_external_loads:
130+
ground_plane = ea.Plane(
131+
plane_origin=np.array([0.0, -base_length * 0.011, 0.0], dtype=np.float64),
132+
plane_normal=normal,
133+
)
134+
sim.append(ground_plane)
133135

134136
for idx in range(n_snakes):
135137
rod = build_rod(
@@ -142,35 +144,36 @@ def build_cpu_sim(
142144
start = snake_start(idx, spacing)
143145
rod.position_collection[...] = rod.position_collection + start[:, None]
144146
sim.append(rod)
145-
sim.add_forcing_to(rod).using(
146-
ea.GravityForces,
147-
acc_gravity=np.array([0.0, gravitational_acc, 0.0], dtype=np.float64),
148-
)
149-
sim.add_forcing_to(rod).using(
150-
ea.MuscleTorques,
151-
base_length=base_length,
152-
b_coeff=b_coeff[:-1],
153-
period=period,
154-
wave_number=2.0 * np.pi / wave_length,
155-
phase_shift=0.0,
156-
rest_lengths=rod.rest_lengths,
157-
ramp_up_time=period,
158-
direction=normal,
159-
with_spline=True,
160-
)
161-
sim.detect_contact_between(rod, ground_plane).using(
162-
ea.RodPlaneContactWithAnisotropicFriction,
163-
k=1.0,
164-
nu=1.0e-6,
165-
slip_velocity_tol=1.0e-8,
166-
static_mu_array=static_mu_array,
167-
kinetic_mu_array=kinetic_mu_array,
168-
)
169-
sim.dampen(rod).using(
170-
ea.AnalyticalLinearDamper,
171-
damping_constant=DEFAULT_DAMPING,
172-
time_step=time_step,
173-
)
147+
if include_external_loads:
148+
sim.add_forcing_to(rod).using(
149+
ea.GravityForces,
150+
acc_gravity=np.array([0.0, gravitational_acc, 0.0], dtype=np.float64),
151+
)
152+
sim.add_forcing_to(rod).using(
153+
ea.MuscleTorques,
154+
base_length=base_length,
155+
b_coeff=b_coeff[:-1],
156+
period=period,
157+
wave_number=2.0 * np.pi / wave_length,
158+
phase_shift=0.0,
159+
rest_lengths=rod.rest_lengths,
160+
ramp_up_time=period,
161+
direction=normal,
162+
with_spline=True,
163+
)
164+
sim.detect_contact_between(rod, ground_plane).using(
165+
ea.RodPlaneContactWithAnisotropicFriction,
166+
k=1.0,
167+
nu=1.0e-6,
168+
slip_velocity_tol=1.0e-8,
169+
static_mu_array=static_mu_array,
170+
kinetic_mu_array=kinetic_mu_array,
171+
)
172+
sim.dampen(rod).using(
173+
ea.AnalyticalLinearDamper,
174+
damping_constant=DEFAULT_DAMPING,
175+
time_step=time_step,
176+
)
174177
rods.append(rod)
175178

176179
sim.finalize()
@@ -190,6 +193,7 @@ def build_jax_sim(
190193
poisson_ratio: float,
191194
gravitational_acc: float,
192195
time_step: float,
196+
include_external_loads: bool = True,
193197
) -> tuple[MultiSnakeJAXSimulator, ea.MemoryBlockCosseratRodJax]:
194198
b_coeff = default_b_coeff()
195199
mu = base_length / (period * period * np.abs(gravitational_acc) * DEFAULT_FROUDE)
@@ -213,28 +217,31 @@ def build_jax_sim(
213217
start = snake_start(idx, spacing)
214218
rod.position_collection[...] = rod.position_collection + start[:, None]
215219
sim.append(rod)
216-
sim.using(rod).operate(
217-
SnakeMuscleTorquesJax,
218-
b_coeff=b_coeff,
219-
period=period,
220-
base_length=base_length,
221-
gravitational_acc=gravitational_acc,
222-
)
223-
sim.using(rod).operate(
224-
SnakePlaneContactJax,
225-
plane_origin=np.array([0.0, -base_length * 0.011, 0.0], dtype=np.float64),
226-
plane_normal=np.array([0.0, 1.0, 0.0], dtype=np.float64),
227-
slip_velocity_tol=1.0e-8,
228-
k=1.0,
229-
nu=1.0e-6,
230-
static_mu_array=static_mu_array,
231-
kinetic_mu_array=kinetic_mu_array,
232-
)
233-
sim.using(rod).operate(
234-
ea.AnalyticalLinearDamperJax,
235-
time_step=np.float64(time_step),
236-
damping_constant=DEFAULT_DAMPING,
237-
)
220+
if include_external_loads:
221+
sim.using(rod).operate(
222+
SnakeMuscleTorquesJax,
223+
b_coeff=b_coeff,
224+
period=period,
225+
base_length=base_length,
226+
gravitational_acc=gravitational_acc,
227+
)
228+
sim.using(rod).operate(
229+
SnakePlaneContactJax,
230+
plane_origin=np.array(
231+
[0.0, -base_length * 0.011, 0.0], dtype=np.float64
232+
),
233+
plane_normal=np.array([0.0, 1.0, 0.0], dtype=np.float64),
234+
slip_velocity_tol=1.0e-8,
235+
k=1.0,
236+
nu=1.0e-6,
237+
static_mu_array=static_mu_array,
238+
kinetic_mu_array=kinetic_mu_array,
239+
)
240+
sim.using(rod).operate(
241+
ea.AnalyticalLinearDamperJax,
242+
time_step=np.float64(time_step),
243+
damping_constant=DEFAULT_DAMPING,
244+
)
238245

239246
sim.finalize()
240247
block = tuple(sim.final_systems())[0]
@@ -348,6 +355,7 @@ def benchmark_config(
348355
n_snakes: int,
349356
n_elem: int,
350357
dt: float,
358+
include_external_loads: bool = True,
351359
) -> dict[str, Any]:
352360
return {
353361
"n_snakes": n_snakes,
@@ -359,4 +367,5 @@ def benchmark_config(
359367
"poisson_ratio": DEFAULT_POISSON_RATIO,
360368
"gravitational_acc": DEFAULT_GRAVITY,
361369
"time_step": dt,
370+
"include_external_loads": include_external_loads,
362371
}

benchmark/jax_snake_io_restart.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,7 @@ def parse_args() -> argparse.Namespace:
4545
parser.add_argument("--n-elem", type=int, default=DEFAULT_N_ELEM)
4646
parser.add_argument("--dt", type=float, default=DEFAULT_DT)
4747
parser.add_argument("--iterations", type=int, default=10)
48+
parser.add_argument("--no-external-loads", action="store_true")
4849
parser.add_argument("--log", type=Path, default=None)
4950
return parser.parse_args()
5051

@@ -60,7 +61,12 @@ def main() -> None:
6061
dtype = np.dtype(np.float32 if args.dtype == "float32" else np.float64)
6162
validate_dtype_for_device(dtype, device)
6263
backend_label = "jax-cpu" if device.platform == "cpu" else f"jax-{device.platform}"
63-
config = benchmark_config(n_snakes=n_snakes, n_elem=args.n_elem, dt=args.dt)
64+
config = benchmark_config(
65+
n_snakes=n_snakes,
66+
n_elem=args.n_elem,
67+
dt=args.dt,
68+
include_external_loads=not args.no_external_loads,
69+
)
6470

6571
with tempfile.TemporaryDirectory(prefix="snake_restart_io_") as tmp_dir:
6672
tmp_path = Path(tmp_dir)
@@ -129,6 +135,7 @@ def _load_jax_state() -> None:
129135
f"n_snakes: {n_snakes}",
130136
f"n_elem: {args.n_elem}",
131137
f"iterations: {args.iterations}",
138+
f"no_external_loads: {args.no_external_loads}",
132139
f"numba_instantiate_avg_seconds: {numba_instantiate_avg:.6f}",
133140
f"numba_save_avg_seconds: {numba_save_avg:.6f}",
134141
f"numba_load_avg_seconds: {numba_load_avg:.6f}",

benchmark/jax_snake_throughput.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,7 @@ def parse_args() -> argparse.Namespace:
4545
parser.add_argument("--steps", type=int, default=DEFAULT_STEPS)
4646
parser.add_argument("--dt", type=float, default=DEFAULT_DT)
4747
parser.add_argument("--warmup-runs", type=int, default=1)
48+
parser.add_argument("--no-external-loads", action="store_true")
4849
parser.add_argument(
4950
"--transfer-guard",
5051
choices=("allow", "log", "disallow", "log_explicit", "disallow_explicit"),
@@ -66,7 +67,12 @@ def main() -> None:
6667
dtype = np.dtype(np.float32 if args.dtype == "float32" else np.float64)
6768
validate_dtype_for_device(dtype, device)
6869
backend_label = "jax-cpu" if device.platform == "cpu" else f"jax-{device.platform}"
69-
config = benchmark_config(n_snakes=n_snakes, n_elem=args.n_elem, dt=args.dt)
70+
config = benchmark_config(
71+
n_snakes=n_snakes,
72+
n_elem=args.n_elem,
73+
dt=args.dt,
74+
include_external_loads=not args.no_external_loads,
75+
)
7076
final_time = np.float64(args.steps * args.dt)
7177

7278
numba_instantiate_start = time.perf_counter()
@@ -140,6 +146,7 @@ def main() -> None:
140146
f"steps: {args.steps}",
141147
f"dt: {args.dt}",
142148
f"warmup_runs: {args.warmup_runs}",
149+
f"no_external_loads: {args.no_external_loads}",
143150
f"transfer_guard: {args.transfer_guard}",
144151
f"numba_instantiate_seconds: {numba_instantiate_elapsed:.6f}",
145152
f"{backend_label}_instantiate_seconds: {jax_instantiate_elapsed:.6f}",

0 commit comments

Comments
 (0)