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