Repository navigation
Add a strictly equivariant model for symmetry-aware RL - #241
vanillaturtlechips wants to merge 5 commits into
Conversation
|
Hi @vanillaturtlechips, Thanks! |
Adds EquivariantMLPModel, which constrains the network so that the actor is equivariant and the critic invariant under a reflection of the robot, by construction rather than through a loss. Extends MLPModel in the same way as CNNModel and RNNModel; PPO and the existing symmetry config are untouched.
The weight projection averages each entry with its independently initialized mirror partner, which halves its variance; across a deep network this made the initial output several times smaller than a plain MLP's. Scale the raw weight so the projected weight starts at the nn.Linear scale. Distributions with a structured MLP output (heteroscedastic Gaussian, Beta) are not supported. Raise a clear NotImplementedError instead of a confusing size-mismatch error. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Running observation statistics are generally not symmetric, so a plain EmpiricalNormalization breaks equivariance. Add SymmetricEmpiricalNormalization, which learns the statistics of each batch together with its mirror image and projects them onto the symmetric subspace after every update, so that normalization commutes with the symmetry exactly. EquivariantMLPModel uses it when obs_normalization=True instead of rejecting the flag. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…ugmentation helper - Rename equivariant_mlp_model.py to mlp_model_equivariant.py - Move EquivariantGaussianDistribution into modules/equivariant.py - Document the feature in the overview, feature list and API reference - Add symmetry_cfg_from_augmentation to derive perm/sign from the symmetry extension's data augmentation function - Fix the citation of the IROS 2024 paper (Su et al.)
88f0d57 to
b73c6c1
Compare
|
Thanks for creating the
Regarding the augmentation function: since it is a signed permutation, the The one part I couldn't make fully clean is that these functions need obs = env.get_observations()
agent_cfg["actor"]["symmetry_cfg"] = symmetry_cfg_from_augmentation(
compute_symmetric_states, env, obs, agent_cfg["obs_groups"]["actor"], env.num_actions
)If you're OK with a small core change, a one-line |
|
Actually, I am not sure anymore if supporting the augmentation function makes sense. I saw that you just convert it to your symmetry config structure, that might make things more confusing than easier. Since you don't really interact with the symmetry extension and already call your classes equivariant, I would rename your configs to equivariance_cfg or something similar. Currently, there are two completely different symmetry_cfg dicts in the repo. |
|
Dropping the connection to the symmetry extension should make things more clear, it would be great if you could also document the configuration of your classes in the configuration documentation. |
|
Comments from my agent:
Questions for the author |
- Remove symmetry_cfg_from_augmentation and rename symmetry_cfg to equivariance_cfg - Document the model and distribution in the configuration guide - Fold the equivariant MLP in as_jit/as_onnx, and create folded layers on the weight's device - Remove SignedPermutation.to() and the group_order argument of regular() - Add future annotations to the tests for Python 3.9 - Fix contributor name
|
Thanks for the review! I agree, keeping it separate from the symmetry extension is clearer.
|
Closes #237.
Adds
EquivariantMLPModel, a model whose network is equivariant by construction under a reflection of the robot, so thatπ(Ms) = Mπ(s)andV(Ms) = V(s)hold exactly (up to floating-point precision) throughout training, instead of being encouraged through data augmentation or a mirror loss.As discussed in the issue, this PR targets the
extrasbranch.Changes
The change is purely additive. PPO, the existing models, and the symmetry extension are untouched.
The
__init__.pyexports only enable the shortclass_name. Without them, the model can also be used through a qualified name ("rsl_rl.models.mlp_model_equivariant:EquivariantMLPModel") viaresolve_callable.Design
EquivariantLinearprojects its weight onto the equivariant subspace on every forward pass,W_eq = ½(W + M_out W M_in)andb_eq = ½(b + M_out b). Gradients flow to the unconstrained parameter. The raw weight is rescaled at init so the projected weight starts at thenn.Linearscale.EquivariantMLPModelsubclassesMLPModeland replaces onlyself.mlp, the same wayCNNModelandRNNModelextend it. Omittingequivariance_cfg["output"]gives the trivial output representation, i.e. an invariant critic.EquivariantGaussianDistributionaverages the std over mirrored action pairs when it's used, so the sampled action distribution is equivariant and not only the mean. The stored parameter and optimizer state are left unchanged.SymmetricEmpiricalNormalizationis used whenobs_normalization=True. Running statistics of real data are generally not symmetric, which would break equivariance, so it learns each batch together with its mirror image and projects the statistics onto the symmetric subspace after every update. This makes normalization commute with the symmetry exactly.fold()turns a trained layer or MLP into plainnn.Linear/nn.Sequentialmodules computing the same function.as_jit()/as_onnx()use it, so the exported policy is a plain MLP without the per-forward projection.Usage
The configuration is documented in the configuration guide.
Testing
Known limitations / possible follow-ups
NotImplementedError.Z2 × Z2(e.g. left-right × front-back for quadrupeds) would generalize the projection from an order-two average to a group average.🤖 Generated with Claude Code