tf.contrib.learn.multi_label_head
Создаёт Head для многоклассовой классификации. (устарело)
tf.contrib.learn.multi_label_head(
n_classes, label_name=None, weight_column_name=None, enable_centered_bias=False,
head_name=None, thresholds=None, metric_class_ids=None, loss_fn=None
)
Многоклассовая классификация обрабатывает случай, когда каждый пример может иметь ноль или более связанных меток из дискретного набора. Это отличается от multi_class_head , у которой ровно одна метка из дискретного набора.
Этот head по умолчанию использует функцию потерь сигмоидальной кросс-энтропии, которая ожидает в качестве входных данных многомерный тензор формы (batch_size, num_classes).
| Аргументы | |
|---|---|
n_classes | Целое число, количество классов, должно быть ≥ 2 |
label_name | Строка, имя ключа в словаре меток. Может быть null, если метка является тензором (модели с одной головой). |
weight_column_name | Строка, определяющая имя столбца признаков, представляющего веса. Используется для уменьшения или увеличения весов примеров во время обучения. Будет умножаться на потерю примера. |
enable_centered_bias | Логическое значение. Если True, estimator будет обучать переменную смещения для каждого класса. Остальная часть структуры модели будет обучать остаток после центрального смещения. |
head_name | Имя head. Если указано, ключи прогнозов, отчётов и метрик будут дополнены суффиксом "/" + head_name , а по умолчанию область переменных будет head_name. |
thresholds | Пороговые значения для оценочных метрик, по умолчанию [.5] |
metric_class_ids | Список идентификаторов классов, для которых необходимо отчитываться по метрикам на каждый класс. Все они должны находиться в диапазоне [0, n_classes). |
loss_fn | Необязательная функция, принимающая (labels, logits, weights) в качестве параметра и возвращающая взвешенную скалярную функцию потерь. weights должно быть необязательным. См. tf.losses |
| Возвращаемое значение | |
|---|---|
Экземпляр Head для многоклассовой классификации. |
| Исключения | |
|---|---|
ValueError | Если n_classes < 2 |
ValueError | Если функция потерь loss_fn не имеет ожидаемой сигнатуры. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/contrib/learn/multi_label_head