Spec-Zone.ru › TensorFlow 1.15

tf.contrib.estimator.boosted_trees_classifier_train_in_memory

Обучает классификатор с усиленными деревьями с использованием набора данных в памяти.

tf.contrib.estimator.boosted_trees_classifier_train_in_memory(
    train_input_fn, feature_columns, model_dir=None,
    n_classes=canned_boosted_trees._HOLD_FOR_MULTI_CLASS_SUPPORT,
    weight_column=None, label_vocabulary=None, n_trees=100, max_depth=6,
    learning_rate=0.1, l1_regularization=0.0, l2_regularization=0.0,
    tree_complexity=0.0, min_node_weight=0.0, config=None, train_hooks=None,
    center_bias=False, pruning_mode='none', quantile_sketch_epsilon=0.01
)

Пример:

bucketized_feature_1 = bucketized_column(
  numeric_column('feature_1'), BUCKET_BOUNDARIES_1)
bucketized_feature_2 = bucketized_column(
  numeric_column('feature_2'), BUCKET_BOUNDARIES_2)

def train_input_fn():
  dataset = create-dataset-from-training-data
  # This is tf.data.Dataset of a tuple of feature dict and label.
  #   e.g. Dataset.zip((Dataset.from_tensors({'f1': f1_array, ...}),
  #                     Dataset.from_tensors(label_array)))
  # The returned Dataset shouldn't be batched.
  # If Dataset repeats, only the first repetition would be used for training.
  return dataset

classifier = boosted_trees_classifier_train_in_memory(
    train_input_fn,
    feature_columns=[bucketized_feature_1, bucketized_feature_2],
    n_trees=100,
    ... <some other params>
)

def input_fn_eval():
  ...
  return dataset

metrics = classifier.evaluate(input_fn=input_fn_eval, steps=10)
Аргументы
train_input_fn функция-обработчик входных данных возвращает набор данных, содержащий одну эпоху неразбитых признаков и меток.
feature_columns Итерируемый объект, содержащий все столбцы признаков, используемые моделью. Все элементы в наборе должны быть экземплярами классов, производных от FeatureColumn.
model_dir Каталог для сохранения параметров модели, графа и т. д. Также можно использовать для загрузки контрольных точек из каталога в оценщик для продолжения обучения ранее сохранённой модели.
n_classes количество классов меток. По умолчанию используется двоичная классификация. Поддержка многоклассовой классификации пока не реализована.
weight_column Строка или объект _NumericColumn, созданный с помощью tf.feature_column.numeric_column, определяющий столбец признака, представляющий веса. Используется для уменьшения или увеличения важности примеров во время обучения. Он будет умножаться на потерю примера. Если это строка, она используется как ключ для извлечения тензора весов из features. Если это объект _NumericColumn, необработанный тензор извлекается по ключу weight_column.key, а затем функция weight_column.normalizer_fn применяется к нему для получения тензора весов.
label_vocabulary Список строк, представляющих возможные значения меток. Если задан, метки должны быть строкового типа и иметь любое значение в label_vocabulary. Если он не задан, это означает, что метки уже закодированы как целые или вещественные числа в диапазоне [0, 1] для n_classes=2 и закодированы как целые числа в {0, 1,..., n_classes-1} для n_classes>2. Также будут ошибки, если словарь не предоставлен, а метки являются строками.
n_trees количество деревьев, которые необходимо создать.
max_depth максимальная глубина дерева, которую необходимо вырастить.
learning_rate параметр уменьшения, который необходимо использовать при добавлении дерева к модели.
l1_regularization множитель регуляризации, применяемый к абсолютным весам листьев дерева.
l2_regularization множитель регуляризации, применяемый к квадратным весам листьев дерева.
tree_complexity коэффициент регуляризации для наказания деревьев с большим количеством листьев.
min_node_weight минимальное значение гессиана, которое должен иметь узел для того, чтобы разбиение было рассмотрено. Значение сравнивается со значением sum(leaf_hessian)/(batch_size * n_batches_per_layer).
config объект RunConfig для настройки параметров выполнения.
train_hooks список экземпляров Hook, которые необходимо передать в estimator.train()
center_bias Требуется ли центрирование смещения. Центрирование смещения относится к первому узлу в самом первом дереве, возвращающему прогноз, который согласуется с исходным распределением меток. Например, для задач регрессии первый узел вернет среднее значение меток. Для задач двоичной классификации он вернет логарифм вероятности метки 1.
pruning_mode одно из значений 'none', 'pre', 'post' для указания отсутствия обрезки, предварительной обрезки (не разбивать узел, если не наблюдается достаточного прироста) и последующей обрезки (построить дерево до максимальной глубины, а затем обрезать ветви с отрицательным приростом). Для предварительной и последующей обрезки НЕОБХОДИМО задать tree_complexity > 0.
quantile_sketch_epsilon число с плавающей точкой от 0 до 1. Граница ошибки для вычисления квантиля. Используется только для столбцов признаков с плавающей точкой, и количество ведер, генерируемых на признак с плавающей точкой, равно 1/quantile_sketch_epsilon.
Возвращает
экземпляр BoostedTreesClassifier , созданный с заданными аргументами и обученный с данными, загруженными в память из input_fn.
Возбуждает
ValueError при указании неправильных аргументов или запросе неподдерживаемых функций.

© 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/estimator/boosted_trees_classifier_train_in_memory

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API