Skip to content

[BugFix] Keep categorical modes one-hot when probabilities tie - #4509

Open
AHMETHAKANBEZIR1 wants to merge 1 commit into
pytorch:mainfrom
AHMETHAKANBEZIR1:fix/one-hot-categorical-tied-mode
Open

AHMETHAKANBEZIR1 wants to merge 1 commit into
pytorch:mainfrom
AHMETHAKANBEZIR1:fix/one-hot-categorical-tied-mode

Conversation

@AHMETHAKANBEZIR1

@AHMETHAKANBEZIR1 AHMETHAKANBEZIR1 commented Sep 30, 2026 •

Copy link
Copy Markdown

Description

Closes #4508 (assigned to AHMETHAKANBEZIR1 through the repository's /assign workflow).

Use argmax followed by one_hot for OneHotCategorical.mode and the dense-mask MaskedOneHotCategorical.mode path. Comparing every value with the maximum previously returned multi-hot vectors when maxima tied, including uniform initial policies. deterministic_sample and OneHotOrdinal inherit the corrected behavior.

The existing sparse-mask mode path remains intact, preserving original action-index remapping and its first-maximum behavior when valid indices are unsorted.

Regression coverage

  • Logits/probs, complete ties, partial ties, and unique maxima, with several batch shapes.
  • Dense masks exclude a larger invalid logit; sparse-mask controls retain unsorted original indices.
  • Ordinal inheritance and deterministic_sample.
  • Two eager-backend full-graph Dynamo checks for the changed unmasked/dense-mask mode paths.

Hard-coded expected action indices provide the regression reference; existing sampling, log-probability and gradient tests are retained. CPU/CUDA parametrization uses the existing get_default_devices convention; only CPU was available locally.

Validation

  • Before the fix, the new non-compile mode matrix had 13 failures and 6 passing sparse-mask controls on main 8f92a50cb1e8543d6d7545605e4e5b399224928b.
  • After the fix: pytest test/test_distributions.py -k 'OneHotCategorical or MaskedCategorical or Ordinal' -q — 162 passed, including the two full-graph eager-backend compile checks.
  • pre-commit run --all-files — all hooks passed with the repository's pinned versions. A separate Python 3.9 tooling environment and PYTHONUTF8=1 were used for Windows compatibility.
  • git diff --check — passed.

Runtime: Windows, Python 3.12.14, PyTorch 2.10.0+cpu, TensorDict 0.14.2. TorchRL Python source is imported from the checkout. The optional C++ extension was not built; its absence warning is reported, and these distribution tests do not use it.

CUDA, Inductor, cudagraphs, full repository tests and the full documentation build were not run locally. The sparse-mask mode branch is unchanged; its existing super-property lookup was not claimed to be fully graph-capturable on this PyTorch version. An initial draft delegating all modes through super().mode caused a Dynamo graph break on the changed paths, so the final fix keeps their direct logits-based computation.

Checklist

  • Raised and claimed an issue for the bug fix.
  • Read CONTRIBUTING.md and CLAUDE.md.
  • Added behavioral tests in the existing test file.
  • Ran focused tests and all-file lint checks.

No new API/signature is introduced. Codex assistance with implementation, investigation and testing is disclosed here and in the commit. The contributor reports having completed the Meta CLA; the repository's automatic verification remains separate.

Co-authored-by: Codex <noreply@openai.com>
@pytorch-bot

pytorch-bot Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

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

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

⚠️ 16 Awaiting Approval

As of commit e6d0036 with merge base 8f92a50 (image):

AWAITING APPROVAL - The following workflows need approval before CI can run:

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

This branch has not been deployed

No deployments
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. distributions Integrations/torch_geometric Integrations Modules

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[BUG] One-hot categorical modes produce invalid actions when maxima are tied

1 participant