sparse_feature_a = sparse_column_with_hash_bucket(...)
sparse_feature_b = sparse_column_with_hash_bucket(...)
sparse_feature_a_emb = embedding_column(sparse_id_column=sparse_feature_a,
...)
sparse_feature_b_emb = embedding_column(sparse_id_column=sparse_feature_b,
...)
estimator = DNNClassifier(
feature_columns=[sparse_feature_a_emb, sparse_feature_b_emb],
hidden_units=[1024, 512, 256])
# Or estimator using the ProximalAdagradOptimizer optimizer with
# regularization.
estimator = DNNClassifier(
feature_columns=[sparse_feature_a_emb, sparse_feature_b_emb],
hidden_units=[1024, 512, 256],
optimizer=tf.compat.v1.train.ProximalAdagradOptimizer(
learning_rate=0.1,
l1_regularization_strength=0.001
))
# Input builders
def input_fn_train: # returns x, y (where y represents label's class index).
pass
estimator.fit(input_fn=input_fn_train)
def input_fn_eval: # returns x, y (where y represents label's class index).
pass
estimator.evaluate(input_fn=input_fn_eval)
def input_fn_predict: # returns x, None
pass
# predict_classes returns class indices.
estimator.predict_classes(input_fn=input_fn_predict)
Если пользователь указывает label_keys в конструкторе, метки должны быть строками из label_keys словаря. Пример:
label_keys = ['label0', 'label1', 'label2']
estimator = DNNClassifier(
feature_columns=[sparse_feature_a_emb, sparse_feature_b_emb],
hidden_units=[1024, 512, 256],
label_keys=label_keys)
def input_fn_train: # returns x, y (where y is one of label_keys).
pass
estimator.fit(input_fn=input_fn_train)
def input_fn_eval: # returns x, y (where y is one of label_keys).
pass
estimator.evaluate(input_fn=input_fn_eval)
def input_fn_predict: # returns x, None
# predict_classes returns one of label_keys.
estimator.predict_classes(input_fn=input_fn_predict)
Входные данные fit и evaluate должны иметь следующие характеристики, в противном случае произойдёт KeyError:
если weight_column_name не None, признак с key=weight_column_name, значение которого является Tensor.
для каждого column в feature_columns:
если column является SparseColumn, признак с key=column.name, значение которого является SparseTensor.
если column является WeightedSparseColumn, два признака: первый с key именем столбца идентификатора, второй с key именем столбца весов. Значение обоих признаков должно быть SparseTensor.
если column является RealValuedColumn, признак с key=column.name, значение которого является Tensor.
Аргументы
hidden_units
Список скрытых узлов на слой. Все слои полностью соединены. Например, [64, 32] означает, что первый слой имеет 64 узла, а второй — 32.
feature_columns
Итерируемый объект, содержащий все столбцы признаков, используемые моделью. Все элементы набора должны быть экземплярами классов, полученных из FeatureColumn.
model_dir
Каталог для сохранения параметров модели, графа и т. д. Это также может быть использовано для загрузки контрольных точек из каталога в оценщик для продолжения обучения ранее сохранённой модели.
n_classes
Количество классов меток. По умолчанию — бинарная классификация. Должно быть больше 1. Примечание: метки классов — это целые числа, представляющие индекс класса (т. е. значения от 0 до n_classes-1). Для произвольных значений меток (например, строковых меток) сначала преобразуйте их в индексы классов.
weight_column_name
Строка, определяющая имя столбца признака, представляющего веса. Используется для уменьшения или увеличения примеров во время обучения. Будет умножено на потерю примера.
optimizer
Экземпляр tf.Optimizer, используемый для обучения модели. Если None, будет использоваться оптимизатор Adagrad.
activation_fn
Функция активации, применяемая к каждому слою. Если None, будет использоваться tf.nn.relu. Также может быть предоставлена строка, содержащая неопределённое имя операции, например, "relu", "tanh" или "sigmoid".
dropout
Если не None, вероятность, что заданная координата будет исключена.
gradient_clip_norm
Число > 0. Если задано, градиенты обрезаются до их глобальной нормы с этим коэффициентом обрезки. Подробнее см. tf.clip_by_global_norm.
enable_centered_bias
Булево значение. Если True, оценщик будет учиться централизованной переменной смещения для каждого класса. Остальная часть структуры модели учится остаточному значению после централизованного смещения.
config
Объект RunConfig для настройки параметров выполнения.
feature_engineering_fn
Функция обработки признаков. Принимает признаки и метки, которые являются результатом input_fn, и возвращает признаки и метки, которые будут переданы в модель.
embedding_lr_multipliers
Необязательно. Словарь от EmbeddingColumn до float множителя. Множитель будет использоваться для умножения на скорость обучения для переменных встраивания.
input_layer_min_slice_size
Необязательно. Минимальный размер фрагмента разделов входного слоя. Если не указано, будет использоваться значение по умолчанию 64 МБ.
label_keys
Необязательный список строк размером [n_classes] для определения словаря меток. Поддерживается только для n_classes > 2.
Исключения
ValueError
Если n_classes < 2.
Атрибуты
config
model_dir
Возвращает путь, по которому процесс оценки будет искать контрольные точки.
Экспортирует граф инференции как SavedModel в указанный каталог.
Аргументы
export_dir_base
Строка, содержащая каталог для записи экспортированного графа и контрольных точек.
serving_input_fn
Функция, которая не принимает аргументов и возвращает InputFnOps.
default_output_alternative_key
Имя головки для обработки, если не указано. Не нужно для моделей с одной головкой.
assets_extra
Словарь, определяющий способ заполнения каталога assets.extra в экспортированном SavedModel. Каждый ключ должен указать целевой путь (включая имя файла) относительно каталога assets.extra. Соответствующее значение указывает полный путь исходного файла для копирования. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
as_text
Флаг, указывающий, следует ли записывать SavedModel протокол в текстовом формате.
checkpoint_path
Путь к контрольной точке для экспорта. Если None (значение по умолчанию), выбирается самая последняя найденная контрольная точка в каталоге модели.
graph_rewrite_specs
Итерируемый объект GraphRewriteSpec. Каждый элемент создаст отдельный MetaGraphDef в экспортированном SavedModel, помеченный и переписанный, как указано. По умолчанию используется один элемент с использованием стандартного тега обслуживания ("serve") и без перезаписи.
Инкрементальное обучение на наборе образцов. (аргументы устарели)
Ожидается, что этот метод будет вызываться несколько раз последовательно на разных или одних и тех же фрагментах набора данных. Это может реализовать итеративное обучение или обучение вне памяти/онлайн-обучение.
Это особенно полезно, когда весь набор данных слишком велик, чтобы поместиться в памяти одновременно. Или когда модель долго сходится, и вы хотите разбить обучение на части.
Аргументы
x
Матрица формы [n_samples, n_features...]. Может быть итератором, возвращающим массивы признаков. Образцы входных данных для обучения модели. Если установлено, input_fn должно быть None.
y
Вектор или матрица [n_samples] или [n_samples, n_outputs]. Может быть итератором, возвращающим массив меток. Значения меток обучения (метки классов в классификации, вещественные числа в регрессии). Если установлено, input_fn должно быть None.
input_fn
Функция ввода. Если установлено, x, y, и batch_size должны быть None.
steps
Количество шагов, на которых нужно обучить модель. Если None, обучать постоянно.
batch_size
Размер мини-пакета для использования на входе, по умолчанию равен первому измерению x. Должно быть None если input_fn предоставлено.
monitors
Список экземпляров подкласса BaseMonitor. Используется для обратных вызовов внутри цикла обучения.
Возвращаемое значение
self, для цепочки.
Исключения
ValueError
Если по крайней мере одно из x и y предоставлено, и input_fn предоставлено.
Возвращает предсказания для заданных признаков. (устаревшие значения аргументов) (устаревшие значения аргументов)
По умолчанию возвращает предсказанные классы. Но это значение по умолчанию скоро будет удалено. Пользователи должны либо передать outputs, либо вызвать метод predict_classes.
Аргументы
x
признаки.
input_fn
Функция ввода. Если установлено, x должен быть None.
batch_size
Переопределение размера пакета по умолчанию.
outputs
список str, имя выходного значения для предсказания. Если None, возвращает классы.
as_iterable
Если True, возвращает итератор, который продолжает генерировать предсказания для каждого примера до тех пор, пока входные данные не будут исчерпаны. Примечание: Входные данные должны завершиться, если вы хотите, чтобы итератор завершился (например, убедитесь, что вы передали num_epochs=1, если вы используете что-то вроде read_batch_features).
Возвращаемое значение
Массив NumPy предсказанных классов с формой размер_пакета. Каждый предсказанный класс представлен его индексом класса (т. е. целым числом от 0 до n_classes-1). Если outputs установлено, возвращает словарь предсказаний.
Возвращает предсказанные классы для заданных признаков. (устаревшие значения аргументов)
Аргументы
x
признаки.
input_fn
Функция ввода. Если установлено, x должен быть None.
batch_size
Переопределение размера пакета по умолчанию.
as_iterable
Если True, возвращает итератор, который продолжает генерировать предсказания для каждого примера до тех пор, пока входные данные не будут исчерпаны. Примечание: Входные данные должны завершиться, если вы хотите, чтобы итератор завершился (например, убедитесь, что вы передали num_epochs=1, если вы используете что-то вроде read_batch_features).
Возвращаемое значение
Массив NumPy предсказанных классов с формой размер_пакета. Каждый предсказанный класс представлен его индексом класса (т. е. целым числом от 0 до n_classes-1).
Возвращает предсказанные вероятности для заданных признаков. (устаревшие значения аргументов)
Аргументы
x
признаки.
input_fn
Функция ввода. Если установлено, x и y должны быть None.
batch_size
Переопределение размера пакета по умолчанию.
as_iterable
Если True, возвращает итератор, который продолжает генерировать предсказания для каждого примера до тех пор, пока входные данные не будут исчерпаны. Примечание: Входные данные должны завершиться, если вы хотите, чтобы итератор завершился (например, убедитесь, что вы передали num_epochs=1, если вы используете что-то вроде read_batch_features).
Метод работает с простыми оценщиками, а также со вложенными объектами (такими как конвейеры). Первые имеют параметры вида <component>__<parameter>, что позволяет обновлять каждый компонент вложенного объекта.