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 accept pt2 archives (.pt2 saved with torch.export.save; preferred) and the deprecated TorchScript (.pt saved with torch.jit.save) and module-bundle ({"model": <network module>, "metadata": {...}} saved with torch.save) formats.

.pt2 and .onnx checkpoints must record num_layers and hidden_size in metadata; only legacy Torch checkpoints may omit them, since their loaded networks expose a live lstm attribute to inspect.

.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, .pt2, .pt, or .pth checkpoint.

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'}#