Skip to content

[Bugfix] Downcast to float32 when Parallel/SerialEnv is on MPS - #3551

Merged
vmoens merged 1 commit into
pytorch:mainfrom
bsprenger:bsprenger/fix/downcast-in-batched-envs
Mar 11, 2026
Merged

vmoens merged 1 commit into
pytorch:mainfrom
bsprenger:bsprenger/fix/downcast-in-batched-envs

Conversation

@bsprenger

Copy link
Copy Markdown
Collaborator

Description

Fix SerialEnv and ParallelEnv crashing with TypeError: Cannot convert a MPS Tensor to float64 dtype when the parent environment is placed on an MPS device and sub-environments produce float64 observations (the default for many Gym/Gymnasium envs).

Motivation and Context

MPS has no hardware support for 64-bit floating-point. Standard envs (e.g. HalfCheetah-v4) return float64 observations by default. As a result, any user who tried to run a batched environment on MPS with a standard Gym env received a hard crash with an opaque TypeError from PyTorch's MPS backend.

How can this happen: the Parallel/SerialEnv have workers which actually run the rollouts. These workers may run the env on CPU, so it can collect float64 data. The error comes when returning that data to the parent Parallel/SerialEnv. To do so, it has to cast it to the parent's device, which can be a different device from the workers. If the parent is on MPS, then an error is raised as it cannot have float64 tensors.

Even if the workers are on CPU in many cases it is better to have the parent be on MPS, as the parent env can be wrapped in a sequence of transforms which are typically more efficient on MPS. More practically, it is just easier to always keep the parent on GPU as a default, and let the worker envs be cpu/gpu as their defaults specify.

Closes: #3550

  • 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:

  • Bug fix (non-breaking change which fixes an issue)

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 Mar 11, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

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

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

❌ 2 New Failures, 4 Cancelled Jobs, 1 Pending, 1 Unrelated Failure

As of commit e2d0cf7 with merge base b6326f7 (image):

NEW FAILURES - The following jobs have failed:

CANCELLED JOBS - The following jobs were cancelled. Please retry:

BROKEN TRUNK - The following job 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 Mar 11, 2026

@vmoens vmoens left a comment

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.

LGTM thanks!

@vmoens
vmoens merged commit a65ff5e into pytorch:main Mar 11, 2026
133 of 141 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

BugFix CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[BUG] SerialEnv and ParallelEnv crash when placed on an MPS device

2 participants