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
при указании неправильных аргументов или запросе неподдерживаемых функций.