Класс Module
Класс Module является базовым классом для всех модулей в TensorFlow.
Он инкапсулирует переменные, другие модули и операции, позволяя создавать
переиспользуемые компоненты, такие как слои и модели.
При наследовании от Module переменные и подмодули, присвоенные атрибутам,
автоматически отслеживаются. Это дает возможность получать список переменных
через variables и обучаемых переменных через trainable_variables.
Класс также поддерживает вызов через метод __call__.
Синтаксис
class MyModule(tf.Module):
def __init__(self, name=None):
super().__init__(name=name)
# создание переменных и подмодулей
Пример
Давайте создадим простой модуль, который хранит переменную и подмодуль, и выведем список его переменных:
import tensorflow as tf
class MyModule(tf.Module):
def __init__(self):
super().__init__()
self.w = tf.Variable(tf.constant([1, 2, 3, 4, 5]), name="w")
self.dense = tf.keras.layers.Dense(2)
m = MyModule()
print([v.name for v in m.variables])
Результат выполнения кода:
['w:0', 'dense/kernel:0', 'dense/bias:0']
Пример
Давайте создадим модуль, который в методе __call__ выполняет операцию
над входными данными, и вызовем его как функцию:
import tensorflow as tf
class AddModule(tf.Module):
def __init__(self):
super().__init__()
self.bias = tf.Variable(10)
def __call__(self, x):
return x + self.bias
m = AddModule()
t = tf.constant([1, 2, 3, 4, 5])
res = m(t)
print(res)
Результат выполнения кода:
tf.Tensor([11 12 13 14 15], shape=(5,), dtype=int32)
Пример
Давайте создадим модуль с двумя подмодулями и получим список всех обучаемых переменных:
import tensorflow as tf
class TwoLayers(tf.Module):
def __init__(self):
super().__init__()
self.d1 = tf.keras.layers.Dense(3)
self.d2 = tf.keras.layers.Dense(1)
m = TwoLayers()
print([v.name for v in m.trainable_variables])
Результат выполнения кода:
['dense/kernel:0', 'dense/bias:0', 'dense_1/kernel:0', 'dense_1/bias:0']
Смотрите также
-
класс
Module,
который является базовым классом для модулей -
метод
__call__,
который позволяет вызывать модуль как функцию -
метод
variables,
который возвращает список переменных модуля -
метод
trainable_variables,
который возвращает список обучаемых переменных модуля