Skip to content

[BugFix] MultiCategorical.to_one_hot supports ndim > 1 nvec - #4517

Closed
kostasrigatos wants to merge 1 commit into
pytorch:mainfrom
kostasrigatos:fix/multicategorical-to-one-hot-ndim
Closed

kostasrigatos wants to merge 1 commit into
pytorch:mainfrom
kostasrigatos:fix/multicategorical-to-one-hot-ndim

Conversation

@kostasrigatos

Copy link
Copy Markdown

Context

This builds on the recent work in #4493 / #4504.

MultiCategorical.to_one_hot_spec() is fixed and it now correctly handles 2D or higher nvec configurations by unbinding and stacking specs:

a = MultiCategorical([[3, 2], [1, 4]])
a.to_one_hot_spec()  # succeeds, returns a Stacked/MultiOneHot spec

This fix is unrelated to #4493/#4504 (those PRs only touch to_one_hot, not to_one_hot_spec), so it looks like this part of the issue was resolved independently at some point and the issue was never updated to reflect it.

The Problem

MultiCategorical.to_one_hot() — the value-encoding method, not the spec-conversion one above — is still broken for ndim > 1 inputs. It throws a TypeError because it loops through the nvec as if it were always 1D:

a = MultiCategorical([[3, 2], [1, 4]])
a.to_one_hot(a.rand())
# TypeError: one_hot(): argument 'num_classes' (position 2) must be int, not Tensor

This is a distinct code path from what #4504 addresses — that PR's to_one_hot fix is scoped to batched/singleton specs, and explicitly notes that the heterogeneous-row case (this issue) is "a different case" and is left untouched. This PR covers that remaining gap.

Proposed Fix

Essentially a mirror of the logic to_one_hot_spec() already uses:

  • When self.ndim > 1, unbind the spec and the value along the appropriate axis.
  • Recursively apply to_one_hot to each resulting pair.
  • If all resulting one-hot tensors have the same width, stack them into a dense tensor.
  • If widths differ (genuinely heterogeneous nvec, e.g. [[1,2],[3,4]]), raise a ValueError — a dense tensor can't represent rows of different width, which is also why to_one_hot_spec() falls back to a lazy Stacked wrapper instead of a flat MultiOneHot in that case.

This also covers the uniform-batched case from #4493/#4504 (e.g. MultiCategorical([3,2], shape=(4,2))) for free, since unbind doesn't distinguish uniform from heterogeneous rows — both go through the same path.

Testing

Added parametrized tests covering: heterogeneous widths (expects ValueError), matching widths (expects successful dense output), 3-level nested nvec, and an extra leading batch dimension beyond self.shape (to catch axis-misalignment bugs).

Full suite: pytest test/test_specs.py — 1816 passed, 7 skipped (CUDA-only). ufmt/flake8 clean.

Future Scope (left out of this fix)

  1. MultiCategorical.to_categorical is currently a no-op (return val) rather than a real one-hot → categorical decoder — so there's no full round-trip for the ndim > 1 case.
  2. Stacked has no to_categorical/to_categorical_spec override, so the lazy spec returned by to_one_hot_spec() for heterogeneous rows can't itself decode a sample back to categorical form.

PR incoming with the fix + tests. @vmoens @Nicholas022400701 — happy to adjust the approach if either of you sees it differently.

AI disclosure: I used Claude in a conversational/pair-programming capacity. Claude was used for design discussion, code review, and provided debugging guidance while writing this patch and its tests, and for drafting this PR description. All code was written by hand - no autonomous coding agent generated or applied any part of the diff. I've read, run, and understood every change myself.

@pytorch-bot

pytorch-bot Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

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

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

❌ 1 New Failure, 15 Unclassified Failures

As of commit 8122451 with merge base 7040ad0 (image):

NEW FAILURE - The following job has failed:

UNCLASSIFIED FAILURES - DrCI could not classify the following jobs because the workflow did not run on the merge base. The failures may be pre-existing on trunk or introduced by this PR:

  • Build Aarch64 Linux Wheels (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Build Linux Wheels (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Build M1 Wheels (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Build Windows Wheels (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Continuous Benchmark (PR) (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Examples Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Generate documentation (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Habitat Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Libs Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • LLM Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Push Binary Nightly (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • SOTA Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Tutorials Tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Unit-tests on Linux (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)
  • Validate Test Partitioning (gh) (this job did not run on the merge base, so DrCI cannot tell whether the failure is pre-existing)

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

@meta-cla

meta-cla Bot commented Oct 3, 2026

Copy link
Copy Markdown

Hi @kostasrigatos!

Thank you for your pull request and welcome to our community.

Action Required

In order to merge any pull request (code, docs, etc.), we require contributors to sign our Contributor License Agreement, and we don't seem to have one on file for you.

Process

In order for us to review and merge your suggested changes, please sign at https://code.facebook.com/cla. If you are contributing on behalf of someone else (eg your employer), the individual CLA may not be sufficient and your employer may need to sign the corporate CLA.

Once the CLA is signed, our tooling will perform checks and validations. Afterwards, the pull request will be tagged with CLA signed. The tagging process may take up to 1 hour after signing. Please give it that time before contacting us about it.

If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks!

@github-actions github-actions Bot added the BugFix label Oct 3, 2026
@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 Oct 3, 2026
@meta-cla

meta-cla Bot commented Oct 3, 2026

Copy link
Copy Markdown

Thank you for signing our Contributor License Agreement. We can now accept your code for this (and any) Meta Open Source project. Thanks!

@bsprenger

Copy link
Copy Markdown
Collaborator

Heterogenous rows were addressed as part of #4504 already.
Please do not write PR descriptions with agents. There is a PR template that you should use in the future, which is auto-populated if you open a PR manually.

@bsprenger bsprenger closed this Oct 8, 2026
@kostasrigatos

kostasrigatos commented Oct 8, 2026 •

Copy link
Copy Markdown
Author

@bsprenger fair point on the template and thanks for mentioning it. I didn't know that PR descriptions should be 100% human-written, my mistake. I will be switching to the proper etiquette henceforth.

Just a brief comment that I think is worth noting and is addressing the statement "Heterogenous rows were addressed as part of #4504 already". This PR was prior to the heterogeneous-row support of #4504. When I opened it, #4504 explicitly left heterogeneous rows untouched, which was the actual wording of the PR itself, and is what my text above was describing.

Small snippet from the original description of #4504: "#2486 (heterogeneous rows of nvec) is a different case and is not touched," with the same ValueError-based design this PR uses. That's the line I was citing in my own description.

Another point to be considered is that there is still a gap if I am not mistaken. A heterogeneous spec with an extra leading batch dimension beyond its shape, still makes to_one_hot still fail on current main:

spec = MultiCategorical([[1, 2], [3, 4]])
val = spec.rand((10,))
spec.to_one_hot(val)
# RuntimeError: Cannot create a nested tensor with a stack dimension other than 0.
# Got a value of shape torch.Size([10, 2, 2]) for a spec of shape torch.Size([2, 2]).

I am happy to open a follow-up PR for this if you find it useful or fold it into the branch of this PR. Whichever you prefer.

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.

2 participants