tf.saved_model.experimental.TrackableResource
Содержит тензор, который может быть захвачен tf.function.
tf.saved_model.experimental.TrackableResource(
device=''
)
TrackableResource наиболее полезен для состоятельных тензоров, требующих инициализации, таких как tf.lookup.StaticHashTable. TrackableResource обнаруживаются путем обхода графа атрибутов объекта, например, во время tf.saved_model.save.
У TrackableResource есть три метода для переопределения:
-
_create_resourceдолжен создать дескриптор тензора ресурса. -
_initializeдолжен инициализировать ресурс, содержащийся вself.resource_handle. -
_destroy_resourceвызывается при уничтоженииTrackableResourceи должен уменьшить счетчик ссылок ресурса. Для большинства ресурсов это должно выполняться с вызовомtf.raw_ops.DestroyResourceOp.
Пример использования:
class DemoResource(tf.saved_model.experimental.TrackableResource):
def __init__(self):
super().__init__()
self._initialize()
def _create_resource(self):
return tf.raw_ops.VarHandleOp(dtype=tf.float32, shape=[2])
def _initialize(self):
tf.raw_ops.AssignVariableOp(
resource=self.resource_handle, value=tf.ones([2]))
def _destroy_resource(self):
tf.raw_ops.DestroyResourceOp(resource=self.resource_handle)
class DemoModule(tf.Module):
def __init__(self):
self.resource = DemoResource()
def increment(self, tensor):
return tensor + tf.raw_ops.ReadVariableOp(
resource=self.resource.resource_handle, dtype=tf.float32)
demo = DemoModule()
demo.increment([5, 1])
<tf.Tensor: shape=(2,), dtype=float32, numpy=array([6., 2.], dtype=float32)>
| Аргументы | |
|---|---|
device | Строка, указывающая необходимую позицию для данного ресурса, например, "CPU", если этот ресурс должен быть создан на устройстве CPU. Пустое значение устройства позволяет пользователю разместить создание ресурса, поэтому, как правило, оно должно быть пустым, если ресурс имеет смысл только на одном устройстве. |
| Атрибуты | |
|---|---|
resource_handle | Возвращает дескриптор ресурса, связанный с этим ресурсом. |
© 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/saved_model/experimental/TrackableResource