[BUG] `MaskedCategorical` missing `mode` and `deterministic_sample` properties.
Open
@vmoens is already working on this.
Since Oct 9, 2024.
bug
- Dominant language
- Python
- Stars
- 3.6k
- Forks
- 487
- Avg merge
- 1d 1h
- Merged PRs (30d)
- 207
Description
Describe the bug
Similar issue as here. The MaskedCategorical distribution is missing the mode and deterministic_sample properties.
Reason and Possible fixes
The MaskedCategorical distribution should define the following additional properties:
@property
def mode(self) -> torch.Tensor:
if hasattr(self, "logits"):
return self.logits.max(-1, keepdim=True)[1]
return self.probs.max(-1, keepdim=True)[1]
@property
def deterministic_sample(self) -> torch.Tensor:
return self.mode
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Assessment
This issue has not been assessed yet.