Функция random.categorical
Функция random.categorical генерирует случайные целочисленные индексы категорий
из категориального распределения. Первым параметром передается тензор логитов
(ненормализованных логарифмов вероятностей) формы [..., num_classes].
Вторым параметром передается целое число сэмплов num_samples.
Третьим необязательным параметром можно указать зерно seed для воспроизводимости.
Функция возвращает тензор целых чисел формы [..., num_samples],
где каждое значение - индекс выбранной категории.
Синтаксис
tf.random.categorical(logits, num_samples, [seed])
Пример
Давайте сгенерируем пять случайных индексов категорий из распределения, заданного логитами:
import tensorflow as tf
tf.random.set_seed(0)
logits = tf.constant([[1.0, 2.0, 3.0]])
res = tf.random.categorical(logits, 5)
print(res)
Результат выполнения кода:
tf.Tensor([[2 2 2 1 2]], shape=(1, 5), dtype=int64)
Пример
Давайте сгенерируем случайные индексы сразу для нескольких распределений, заданных двумерным тензором логитов:
import tensorflow as tf
tf.random.set_seed(0)
logits = tf.constant([[1.0, 2.0, 3.0], [3.0, 2.0, 1.0]])
res = tf.random.categorical(logits, 4)
print(res)
Результат выполнения кода:
tf.Tensor(
[[2 2 2 1]
[0 0 1 0]], shape=(2, 4), dtype=int64)
Пример
Давайте передадим одинаковые логиты, чтобы получить равномерное распределение по трем категориям:
import tensorflow as tf
tf.random.set_seed(0)
logits = tf.constant([[0.0, 0.0, 0.0]])
res = tf.random.categorical(logits, 6)
print(res)
Результат выполнения кода:
tf.Tensor([[2 2 2 1 2 0]], shape=(1, 6), dtype=int64)
Смотрите также
-
функцию
stateless_categorical,
которая генерирует категории без сохранения состояния -
функцию
uniform,
которая генерирует случайные значения из равномерного распределения -
функцию
normal,
которая генерирует случайные значения из нормального распределения -
функцию
set_seed,
которая задает зерно для генератора случайных чисел