tf.tpu.experimental.embedding.serving_embedding_lookup
Применение стандартных операций поиска с конфигурациями tf.tpu.experimental.embedding.
tf.tpu.experimental.embedding.serving_embedding_lookup(
inputs: Any,
weights: Optional[Any],
tables: Dict[tf.tpu.experimental.embedding.TableConfig, tf.Variable],
feature_config: Union[tf.tpu.experimental.embedding.FeatureConfig, Iterable]
) -> Any
Эта функция — утилита, которая позволяет использовать объекты конфигурации tf.tpu.experimental.embedding со стандартными функциями поиска. Это можно использовать при экспорте модели, которая использует tf.tpu.experimental.embedding.TPUEmbedding для обслуживания на CPU. В частности, tf.tpu.experimental.embedding.TPUEmbedding поддерживает поиск только на TPUs и не должна быть частью вашей служебной графы.
Обратите внимание, что специфичные для TPU параметры (например, max_sequence_length) в объектах конфигурации будут игнорироваться.
В следующем примере мы берем обученную модель (см. документацию для tf.tpu.experimental.embedding.TPUEmbedding для контекста) и создаём сохранённую модель со служебной функцией, которая выполнит поиск встраивания и передаст результаты в вашу модель:
model = model_fn(...)
embedding = tf.tpu.experimental.embedding.TPUEmbedding(
feature_config=feature_config,
batch_size=1024,
optimizer=tf.tpu.experimental.embedding.SGD(0.1))
checkpoint = tf.train.Checkpoint(model=model, embedding=embedding)
checkpoint.restore(...)
@tf.function(input_signature=[{'feature_one': tf.TensorSpec(...),
'feature_two': tf.TensorSpec(...),
'feature_three': tf.TensorSpec(...)}])
def serve_tensors(embedding_features):
embedded_features = tf.tpu.experimental.embedding.serving_embedding_lookup(
embedding_features, None, embedding.embedding_tables,
feature_config)
return model(embedded_features)
model.embedding_api = embedding
tf.saved_model.save(model,
export_dir=...,
signatures={'serving_default': serve_tensors})
Примечание: Важно назначить объект API встраивания члену вашей модели, так какtf.saved_model.saveподдерживает сохранение переменных только как одинTrackableобъект. Поскольку веса модели находятся вmodel, а таблица встраивания управляетсяembedding, мы назначаемembeddingв атрибутmodel, чтобы tf.saved_model.save мог найти переменные встраивания.
Примечание: Та же функцияserve_tensorsи вызовtf.saved_model.saveбудут работать непосредственно из обучения.
| Аргументы | |
|---|---|
inputs | вложенная структура тензоров, разреженных тензоров или разрозненных тензоров. |
weights | вложенная структура тензоров, разреженных тензоров или разрозненных тензоров или None для отсутствия весов. Если не None, структура должна соответствовать структуре входных данных, но допускается наличие None. |
tables | словарь, сопоставляющий объекты TableConfig с переменными. |
feature_config | вложенная структура объектов FeatureConfig с той же структурой, что и входные данные. |
| Возвращает | |
|---|---|
| Вложенная структура тензоров с той же структурой, что и входные данные. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/tpu/experimental/embedding/serving_embedding_lookup