newton.actuators.DriveNeuralLSTM#
- class newton.actuators.DriveNeuralLSTM(model_path)[source]#
Bases:
DriveBaseLSTM-based neural network actuator drive.
Uses a pre-trained LSTM network to compute joint effort from position error and joint velocity. Hidden and cell state are maintained across timesteps.
Torch checkpoints use the Torch backend and preserve the Torch state interface. They accept pt2 archives (
.pt2saved withtorch.export.save; preferred) and the deprecated TorchScript (.ptsaved withtorch.jit.save) and module-bundle ({"model": <network module>, "metadata": {...}}saved withtorch.save) formats..pt2and.onnxcheckpoints must recordnum_layersandhidden_sizein metadata; only legacy Torch checkpoints may omit them, since their loaded networks expose a livelstmattribute to inspect..onnxcheckpoints use Warp-NN. The exported ONNX model must have three inputs (input, initial hidden, and initial cell) and three graph outputs (effort, hidden output, and cell output). Metadata properties map those names to drive roles.ONNX checkpoints support the implicit effort mode through a per-step linearization of the network; Torch checkpoints do not and must use the explicit mode.
- evaluate_force(q, qd, target_q, target_qd, feedforward, params, i)#
The network enters the general implicit solve as a per-step-linearized in-kernel law (see
prepare_implicit()), like any other drive.
- classmethod resolve_arguments(args)#
- __init__(model_path)#
Initialize the LSTM drive from a checkpoint file.
- Parameters:
model_path (str) – Path to the
.onnx,.pt2,.pt, or.pthcheckpoint.
- bind_params()#
Linearization pack
[tau0, a, b, q0, qd0];Noneif implicit unsupported.The pack is allocated in
finalize()and rewritten in place each step byprepare_implicit().Nonefor the Torch backend.
- compute(positions, velocities, target_pos, target_vel, feedforward, pos_indices, vel_indices, target_pos_indices, target_vel_indices, forces, state, dt, device=None)#
- finalize(device, num_actuators)#
- is_graphable()#
- is_stateful()#
- prepare_implicit(positions, velocities, target_pos, target_vel, pos_indices, vel_indices, target_pos_indices, target_vel_indices, drive_state, dt, inv_mass=None, device=None)#
Refresh the linearization of the network about the current state.
One forward + autodiff backward at the current per-slot state, with the incoming hidden/cell state held fixed, gives
tau0, d(tau)/dq, d(tau)/dqd, packed as[tau0, a, b, q0, qd0]intobind_params(). The forward also advances hidden/cell forupdate_state().
- state(num_actuators, device)#
- update_state(current_state, next_state)#