Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
320 of 824 menu
◀ ▶

Метод call класса Module

Метод call класса Module задает вычисления, которые выполняются при прямом проходе через модуль. Когда вы вызываете экземпляр модуля как функцию, TensorFlow автоматически обращается к этому методу. Первым параметром метод принимает входные данные, последующими параметрами могут быть дополнительные аргументы, необходимые для вычислений.

Синтаксис

class MyModule(tf.Module): def __call__(self, inputs, *args, **kwargs): return self.call(inputs, *args, **kwargs) def call(self, inputs, *args, **kwargs): # logic return outputs

Пример

Давайте создадим простой модуль, который умножает входной тензор на заданный коэффициент, и вызовем его:

import tensorflow as tf class ScaleModule(tf.Module): def __init__(self, factor): super().__init__() self.factor = factor def __call__(self, inputs): return self.call(inputs) def call(self, inputs): return inputs * self.factor module = ScaleModule(2) t = tf.constant([1, 2, 3, 4, 5]) res = module(t) print(res)

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

tf.Tensor([ 2 4 6 8 10], shape=(5,), dtype=int32)

Пример

Давайте создадим модуль, который принимает дополнительные аргументы в методе call, и передадим их при вызове:

import tensorflow as tf class AddModule(tf.Module): def __init__(self): super().__init__() def __call__(self, inputs, value): return self.call(inputs, value) def call(self, inputs, value): return inputs + value module = AddModule() t = tf.constant([1, 2, 3, 4, 5]) res = module(t, 10) print(res)

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

tf.Tensor([11 12 13 14 15], shape=(5,), dtype=int32)

Пример

Давайте создадим модуль с обучаемой переменной и используем ее в методе call:

import tensorflow as tf class DenseModule(tf.Module): def __init__(self): super().__init__() self.w = tf.Variable(tf.constant([1, 2, 3])) def __call__(self, inputs): return self.call(inputs) def call(self, inputs): return inputs * self.w module = DenseModule() t = tf.constant([1, 2, 3, 4, 5]) res = module(t) print(res)

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

tf.Tensor([ 1 4 9 16 25], shape=(5,), dtype=int32)

Смотрите также

  • класс Module,
    который является базовым классом для модулей
  • метод __call__,
    который вызывает метод call при обращении к модулю
  • метод variables,
    который возвращает список переменных модуля
  • метод trainable_variables,
    который возвращает список обучаемых переменных модуля
← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить