tf.contrib.learn.multi_class_head
Создаёт Head для многоклассовой классификации с одной меткой. (устарело)
tf.contrib.learn.multi_class_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,
label_keys=None
)
Этот головной модуль использует потерю софтмакс перекрёстной энтропии.
Этот головной модуль ожидает целые метки, указывающие индекс класса. Но если label_keys указано, то метки должны быть строками из этого словаря, а предсказанные классы будут строками из того же словаря.
| Аргументы | |
|---|---|
n_classes | Целое число, количество классов, должно быть ≥ 2 |
label_name | Строка, имя ключа в словаре меток. Может быть null, если метка — это тензор (модели с одним головным модулем). |
weight_column_name | Строка, определяющая имя столбца признака, представляющего веса. Используется для понижения или повышения весов примеров во время обучения. Будет умножаться на потерю примера. |
enable_centered_bias | Булево значение. Если True, оценочный модуль выучит сдвинутую переменную смещения для каждого класса. Остальная часть структуры модели выучит остаток после сдвинутого смещения. |
head_name | Имя головного модуля. Если указано, ключи прогнозов, отчётов и метрик будут дополнены суффиксом "/" + head_name, а стандартный объём переменных будет head_name. |
thresholds | Пороговые значения для метрик оценки, по умолчанию равны [.5] |
metric_class_ids | Список идентификаторов классов, для которых необходимо отчитываться о метриках по классу. Все они должны быть в диапазоне [0, n_classes). Недействительно, если n_classes равно 2. |
loss_fn | Необязательная функция, которая принимает (labels, logits, weights) в качестве параметра и возвращает взвешенную скалярную потерю. weights должно быть необязательным. См. tf.losses |
label_keys | Необязательный список строк размером [n_classes], определяющий словарь меток. Поддерживается только для n_classes > 2. |
| Возвращаемое значение | |
|---|---|
Экземпляр Head для многоклассовой классификации. |
| Исключения | |
|---|---|
ValueError | если n_classes < 2. |
ValueError | Если metric_class_ids указан, когда n_classes равен 2. |
ValueError | Если len(label_keys) != n_classes. |
© 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_class_head