Spec-Zone.ru › TensorFlow

tf.train.experimental.ShardingCallback

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

Просмотр псевдонимов

Псевдонимы совместимости для миграции

Дополнительные сведения см. в руководстве по миграции.

tf.compat.v1.train.experimental.ShardingCallback

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

Обратный вызов принимает список tf.train.experimental.ShardableTensor в качестве входных данных (а также любые kwargs, определённые подклассом tf.train.experimental.ShardingCallback), и организует входные тензоры в различные фрагменты. Тензоры сначала организуются по задаче устройства (см. tf.DeviceSpec), затем обратный вызов вызывается для каждой коллекции тензоров.

При создании пользовательского обратного вызова следует учитывать несколько ограничений:

  • Тензоры не должны удаляться из контрольной точки.
  • Тензоры не должны быть преобразованы.
  • Типы тензоров не должны изменяться.
  • Тензоры в пределах фрагмента должны принадлежать одной задаче. Проверки валидности будут выполнены после выполнения функции обратного вызова, чтобы убедиться, что эти ограничения не нарушены.

Вот пример простого пользовательского обратного вызова:

# Place all tensors in a single shard.
class AllInOnePolicy(tf.train.experimental.ShardingCallback):
  @property
  def description(self):
    return "Place all tensors in a single shard."

  def __call__(self, shardable_tensors):
    tensors = {}
    for shardable_tensor in shardable_tensors:
      tensor = shardable_tensor.tensor_save_spec.tensor
      checkpoint_key = shardable_tensor.checkpoint_key
      slice_spec = shardable_tensor.slice_spec

      tensors.set_default(checkpoint_key, {})[slice_spec] = tensor
    return [tensors]

ckpt.save(
    "path",
    options=tf.train.CheckpointOptions(
        experimental_sharding_callback=AllInOnePolicy()))

Атрибут description используется для идентификации обратного вызова и для помощи в отладке во время сохранения и восстановления.

Для приема kwargs просто определите конструктор и передайте их:

class ParameterPolicy(tf.train.experimental.ShardingCallback):
  def __init__(self, custom_param):
    self.custom_param = custom_param
  ...

ckpt.save(
    "path",
    options=tf.train.CheckpointOptions(
        experimental_sharding_callback=ParameterPolicy(custom_param=...)))
Атрибуты
description

Методы

__call__

Просмотреть исходный код

@abc.abstractmethod
__call__(
    shardable_tensors: Sequence[tf.train.experimental.ShardableTensor]
) -> Sequence[TensorSliceDict]

Вызов self как функции.

© 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/api_docs/python/tf/train/experimental/ShardingCallback

Spec-Zone.ru

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