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
regressor = boosted_trees_regressor_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 = regressor.evaluate(input_fn=input_fn_eval, steps=10)
Аргументы
train_input_fn
функция ввода возвращает набор данных, содержащий один эпох неразбитых признаков и меток.
feature_columns
Итерируемый объект, содержащий все столбцы признаков, используемые моделью. Все элементы набора должны быть экземплярами классов, производных от FeatureColumn.
model_dir
Директория для сохранения параметров модели, графа и т. д. Ее также можно использовать для загрузки контрольных точек из директории в оценщик для продолжения обучения ранее сохраненной модели.
label_dimension
Количество целевых значений регрессии на пример. Многомерная поддержка пока не реализована.
weight_column
Строка или _NumericColumn, созданная с помощью tf.feature_column.numeric_column, определяющая столбец признака, представляющий веса. Используется для уменьшения или увеличения примеров во время обучения. Он будет умножен на потерю примера. Если это строка, она используется в качестве ключа для извлечения тензора весов из features. Если это _NumericColumn, исходный тензор извлекается по ключу weight_column.key, затем weight_column.normalizer_fn применяется к нему для получения тензора весов.
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
при указании неправильных аргументов или запросе неподдерживаемых функций.