Skip to content

[Feature] Add TD-MPC2 Q-network ensemble - #4453

Merged
bsprenger merged 1 commit into
pytorch:mainfrom
bsprenger:feat/tdmpc2-q-ensemble
Sep 29, 2026
Merged

bsprenger merged 1 commit into
pytorch:mainfrom
bsprenger:feat/tdmpc2-q-ensemble

Conversation

@bsprenger

@bsprenger bsprenger commented Sep 21, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Note

This PR stacks on #4452 . I cannot find a good way to compare the diff against that base, since it is in my own fork.
Once it is merged, I will re-target this one to make the diff easier to read.

This PR introduces the TD-MPC2 Q ensemble, which is a key part of the algorithm. The Q-ensemble consists of $N$ parallel Q-functions to compute TD targets during training of the model and to be used during the planning stage of the algorithm.

Motivation and Context

This is the 4th PR in the chain intended to address #4448 . The Q-ensemble is described in depth in the original paper and associated implementation.

  • I have raised an issue to propose this change (required for new features and bug fixes)

Types of changes

  • New feature (non-breaking change which adds core functionality)

Checklist

Go over all the following points, and put an x in all the boxes that apply.
If you are unsure about any of these, don't hesitate to ask. We are here to help!

  • I have read the CONTRIBUTION guide (required)
  • My change requires a change to the documentation.
  • I have updated the tests accordingly (required for a bug fix or a new feature).
  • I have updated the documentation accordingly.

@pytorch-bot

pytorch-bot Bot commented Sep 21, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/rl/4453

Note: Links to docs will display an error until the docs builds have been completed.

✅ You can merge normally! (2 Unrelated Failures)

As of commit d6e7f5e with merge base 71a4143 (image):

BROKEN TRUNK - The following jobs failed but were present on the merge base:

👉 Rebase onto the `viable/strict` branch to avoid these failures

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Sep 21, 2026
@github-actions github-actions Bot added Feature New feature Documentation Improvements or additions to documentation Modules Trainers labels Sep 21, 2026
Comment thread torchrl/trainers/algorithms/tdmpc2.py Outdated
Comment on lines +189 to +190
num_bins: Number of categorical bins. ``0`` uses an identity output and
``1`` uses a symmetric exponential output.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Advertised num_bins modes are rejected.

The Q-ensemble documentation says 0 provides identity regression and 1 provides symlog regression, but the constructor rejects both. The config factories and loss also only support categorical mode.

Either implement the documented/reference behavior or explicitly narrow the API and documentation to num_bins > 1.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I updated the documented behaviour. Thanks!

@bsprenger
bsprenger force-pushed the feat/tdmpc2-q-ensemble branch 2 times, most recently from cc85d9a to 1cdfd84 Compare September 22, 2026 05:50
@bsprenger
bsprenger force-pushed the feat/tdmpc2-q-ensemble branch 3 times, most recently from cafff05 to 620080e Compare September 22, 2026 07:42
@bsprenger
bsprenger force-pushed the feat/tdmpc2-q-ensemble branch from 620080e to d6e7f5e Compare September 28, 2026 21:25
@bsprenger

Copy link
Copy Markdown
Collaborator Author

@theap06 I've rebased this one after merging its base branch, looking for a re-review / CI, thanks!

@theap06 theap06 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM!

@bsprenger
bsprenger merged commit 8a53f8c into pytorch:main Sep 29, 2026
120 of 122 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. Documentation Improvements or additions to documentation Feature New feature Integrations/torch_geometric Integrations Modules Trainers

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants