newton.actuators.DriveNeuralLSTM#

class newton.actuators.DriveNeuralLSTM(model_path)[source]#

Bases: DriveBase

LSTM-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 must be pt2 archives saved with torch.export.save.

Checkpoints must record num_layers and hidden_size in metadata.

.onnx checkpoints 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 checkpoint or the pt2 archive (.pt2, .pt, or .pth).

bind_params()#

Linearization pack [tau0, a, b, q0, qd0]; None if implicit unsupported.

The pack is allocated in finalize() and rewritten in place each step by prepare_implicit(). None for 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] into bind_params(). The forward also advances hidden/cell for update_state().

state(num_actuators, device)#
update_state(current_state, next_state)#
SHARED_PARAMS: ClassVar[set[str]] = {'model_path'}#