tf.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