Метод __call__ класса Module
Метод __call__ класса Module делает экземпляр модуля вызываемым,
как обычную функцию. Именно этот метод стоит за конструкцией
model(x). Внутри он вызывает метод call, который вы
переопределяете в своем классе. Первым параметром в __call__
передаются входные данные, а все остальные позиционные и именованные
аргументы прокидываются в call. Перед вызовом call модуль
автоматически строит свои переменные и переводит их в режим отслеживания,
что избавляет от ручного управления состоянием.
Метод также управляет обучением: в режиме training=True слои
вроде Dropout и BatchNormalization ведут себя иначе.
Возвращаемое значение - это результат работы call, чаще всего
тензор или словарь тензоров.
Синтаксис
module(*args, **kwargs)
Пример
Давайте создадим простой модуль и вызовем его как функцию:
import tensorflow as tf
class Linear(tf.keras.Module):
def __init__(self):
super().__init__()
self.w = tf.Variable(2.0)
self.b = tf.Variable(1.0)
def call(self, x):
return self.w * x + self.b
model = Linear()
res = model(tf.constant([1, 2, 3, 4, 5]))
print(res)
Результат выполнения кода:
tf.Tensor([3. 5. 7. 9. 11.], shape=(5,), dtype=float32)
Пример
Давайте передадим в модуль дополнительный аргумент и посмотрим,
как он прокидывается в метод call:
import tensorflow as tf
class Scaler(tf.keras.Module):
def __init__(self):
super().__init__()
self.w = tf.Variable(1.0)
def call(self, x, scale):
return self.w * x * scale
model = Scaler()
res = model(tf.constant([1, 2, 3, 4, 5]), scale=10)
print(res)
Результат выполнения кода:
tf.Tensor([10. 20. 30. 40. 50.], shape=(5,), dtype=float32)
Пример
Давайте проверим, что вызов модуля автоматически регистрирует
переменные в trainable_variables:
Результат выполнения кода:
tf.Tensor([3. 5. 7. 9. 11.], shape=(5,), dtype=float32)
[<tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>]
Смотрите также
-
класс
Module,
который является базовым классом для всех модулей -
метод
call,
который переопределяется для описания прямого прохода -
метод
variables,
который возвращает список переменных модуля -
метод
trainable_variables,
который возвращает список обучаемых переменных модуля