Repository navigation
[BugFix] MultiCategorical.to_one_hot supports ndim > 1 nvec - #4517
kostasrigatos wants to merge 1 commit into
Conversation
🔗 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 FailuresAs of commit 8122451 with merge base 7040ad0 ( 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:
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
|
Hi @kostasrigatos! Thank you for your pull request and welcome to our community. Action RequiredIn 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. ProcessIn 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 If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks! |
|
Thank you for signing our Contributor License Agreement. We can now accept your code for this (and any) Meta Open Source project. Thanks! |
|
Heterogenous rows were addressed as part of #4504 already. |
|
@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. |
Context
This builds on the recent work in #4493 / #4504.
MultiCategorical.to_one_hot_spec()is fixed and it now correctly handles 2D or highernvecconfigurations by unbinding and stacking specs:This fix is unrelated to #4493/#4504 (those PRs only touch
to_one_hot, notto_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 forndim > 1inputs. It throws aTypeErrorbecause it loops through thenvecas if it were always 1D:This is a distinct code path from what #4504 addresses — that PR's
to_one_hotfix 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:self.ndim > 1, unbind the spec and the value along the appropriate axis.to_one_hotto each resulting pair.nvec, e.g.[[1,2],[3,4]]), raise aValueError— a dense tensor can't represent rows of different width, which is also whyto_one_hot_spec()falls back to a lazyStackedwrapper instead of a flatMultiOneHotin that case.This also covers the uniform-batched case from #4493/#4504 (e.g.
MultiCategorical([3,2], shape=(4,2))) for free, sinceunbinddoesn'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 nestednvec, and an extra leading batch dimension beyondself.shape(to catch axis-misalignment bugs).Full suite:
pytest test/test_specs.py— 1816 passed, 7 skipped (CUDA-only).ufmt/flake8clean.Future Scope (left out of this fix)
MultiCategorical.to_categoricalis currently a no-op (return val) rather than a real one-hot → categorical decoder — so there's no full round-trip for thendim > 1case.Stackedhas noto_categorical/to_categorical_specoverride, so the lazy spec returned byto_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.