Skip to content

[Feature] Add TD-MPC2 loss module - #4454

Merged
bsprenger merged 2 commits into
pytorch:mainfrom
bsprenger:feat/tdmpc2-loss
Sep 30, 2026
Merged

bsprenger merged 2 commits into
pytorch:mainfrom
bsprenger:feat/tdmpc2-loss

Conversation

@bsprenger

Copy link
Copy Markdown
Collaborator

Description

Note

This PR stacks on #4453 . 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 core loss objective for the TD-MPC2 algorithm, porting over the logic into the trainer framework.

Motivation and Context

5th PR in the chain to address #4448 . The design of the loss is 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

What types of changes does your code introduce? Remove all that do not apply:

  • 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/4454

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 5842fcf with merge base 8a53f8c (image):

FLAKY - The following job failed but was likely due to flakiness present on trunk:

BROKEN TRUNK - The following job failed but was 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 and removed Feature New feature labels Sep 21, 2026

@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.

Overall looks pretty good. Just one nitpick: could you add a test case in tdmpc2 that checks via public api as the current tests cover shapes and the private methods?

@bsprenger

Copy link
Copy Markdown
Collaborator Author

Updated to test the public API for model_loss and actor_loss_from_latents. Thanks for the feedback!

@bsprenger
bsprenger requested a review from theap06 September 29, 2026 18:53
@bsprenger
bsprenger merged commit 8f92a50 into pytorch:main Sep 30, 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 Objectives Trainers

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants