TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

OneHotCategoricalStraightThrough

class tensorplay.distributions.OneHotCategoricalStraightThrough(probs: Tensor | None = None, logits: Tensor | None = None, validate_args: bool | None = None)[source]

Creates a reparameterizable OneHotCategorical distribution based on the straight- through gradient estimator from [1].

[1] Estimating or Propagating Gradients Through Stochastic Neurons for Conditional Computation (Bengio et al., 2013)

property batch_shape: Size

Returns the shape over which parameters are batched.

cdf(value: Tensor) → Tensor

Returns the cumulative density/mass function evaluated at value.

Parameters:

value (Tensor)

property event_shape: Size

Returns the shape of a single sample (without batching).

icdf(value: Tensor) → Tensor

Returns the inverse cumulative density/mass function evaluated at value.

Parameters:

value (Tensor)

perplexity() → Tensor

Returns perplexity of distribution, batched over batch_shape.

Returns:

Tensor of shape batch_shape.

sample_n(n: int) → Tensor

Generates n samples or n batches of samples if the distribution parameters are batched.

static set_default_validate_args(value: bool) → None

Sets whether validation is enabled or disabled.

The default behavior mimics Python’s assert statement: validation is on by default, but is disabled if Python is run in optimized mode (via python -O). Validation may be expensive, so you may want to disable it once a model is working.

Parameters:

value (bool) – Whether to enable validation.

property stddev: Tensor

Returns the standard deviation of the distribution.

On this page

Ask DeepWiki