Download flow_matching/utils/categorical_sampler.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 615 Bytes
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/flow_matching/utils/categorical_sampler.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/flow_matching/utils/categorical_sampler.py
-
curl -L -o categorical_sampler.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/flow_matching/utils/categorical_sampler.py
615 Bytes
| # Copyright (c) Meta Platforms, Inc. and affiliates. | |
| # All rights reserved. | |
| # | |
| # This source code is licensed under the CC-by-NC license found in the | |
| # LICENSE file in the root directory of this source tree. | |
| import torch | |
| from torch import Tensor | |
| def categorical(probs: Tensor) -> Tensor: | |
| r"""Categorical sampler according to weights in the last dimension of ``probs`` using :func:`torch.multinomial`. | |
| Args: | |
| probs (Tensor): probabilities. | |
| Returns: | |
| Tensor: Samples. | |
| """ | |
| return torch.multinomial(probs.flatten(0, -2), 1, replacement=True).view( | |
| *probs.shape[:-1] | |
| ) | |