РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
321 of 824 menu

Метод __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:

<+python+> 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) print(model.trainable_variables) <-python+>

Результат выполнения кода:

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,
    который возвращает список обучаемых переменных модуля
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить