Spec-Zone.ru › TensorFlow

tf.keras.layers.FlaxLayer

Слой Keras, который оборачивает модуль Flax.

Наследуется от: JaxLayer, Layer, Operation

tf.keras.layers.FlaxLayer(
    module, method=None, variables=None, **kwargs
)

Этот слой позволяет использовать компоненты Flax в виде экземпляров flax.linen.Module в Keras, когда JAX используется в качестве бэкенда для Keras.

Метод модуля для использования в прямом проходе можно указать с помощью аргумента method и он равен __call__ по умолчанию. Этот метод должен принимать следующие аргументы с этими точными именами:

  • self, если метод связан с модулем, что соответствует значению по умолчанию __call__, и module в противном случае, чтобы передать модуль.
  • inputs: входные данные модели, массив JAX или PyTree массивов.
  • training (необязательно): аргумент, указывающий, находимся ли мы в режиме обучения или вывода, True передается в режиме обучения.

FlaxLayer автоматически обрабатывает состояние, не участвующее в обучении, вашей модели и необходимые RNG. Обратите внимание, что параметр mutable flax.linen.Module.apply() установлен в значение DenyList(["params"]), следовательно, предполагается, что все переменные, не входящие в коллекцию «params», являются весами, не участвующими в обучении.

В этом примере показано, как создать FlaxLayer из Flax Module с методом по умолчанию __call__ и без аргумента обучения:

class MyFlaxModule(flax.linen.Module):
    @flax.linen.compact
    def __call__(self, inputs):
        x = inputs
        x = flax.linen.Conv(features=32, kernel_size=(3, 3))(x)
        x = flax.linen.relu(x)
        x = flax.linen.avg_pool(x, window_shape=(2, 2), strides=(2, 2))
        x = x.reshape((x.shape[0], -1))  # flatten
        x = flax.linen.Dense(features=200)(x)
        x = flax.linen.relu(x)
        x = flax.linen.Dense(features=10)(x)
        x = flax.linen.softmax(x)
        return x

flax_module = MyFlaxModule()
keras_layer = FlaxLayer(flax_module)

В этом примере показано, как обернуть метод модуля, чтобы он соответствовал требуемому подписному значению. Это позволяет иметь несколько аргументов ввода и аргумент обучения, имеющий другое имя и значения. Кроме того, здесь показано, как использовать функцию, которая не связана с модулем.

class MyFlaxModule(flax.linen.Module):
    @flax.linen.compact
    def forward(self, input1, input2, deterministic):
        ...
        return outputs

def my_flax_module_wrapper(module, inputs, training):
    input1, input2 = inputs
    return module.forward(input1, input2, not training)

flax_module = MyFlaxModule()
keras_layer = FlaxLayer(
    module=flax_module,
    method=my_flax_module_wrapper,
)
Аргументы
module Экземпляр flax.linen.Module или подкласс.
method Метод вызова модели. Как правило, это метод в Module. Если не указано, используется метод __call__. method также может быть функцией, не определенной в Module, в этом случае она должна принимать Module в качестве первого аргумента. Она используется для Module.init и Module.apply. Подробности документированы в аргументе method flax.linen.Module.apply().
variables Словарь, содержащий все переменные модуля в том же формате, что и возвращаемый flax.linen.Module.init(). Он должен содержать ключ «params» и, при необходимости, другие ключи для коллекций переменных для состояния, не участвующего в обучении. Это позволяет передавать обученные параметры и состояние, не участвующее в обучении, или управлять инициализацией. Если None передается, функция init модуля вызывается во время построения для инициализации переменных модели.
Атрибуты
input Получает тензор(ы) ввода символической операции.

Возвращает только тензор(ы), соответствующие первому вызову операции.

output Получает тензор(ы) вывода слоя.

Возвращает только тензор(ы), соответствующие первому вызову операции.

Методы

from_config

Просмотреть исходный код

@classmethod
from_config(
    config
)

Создает слой из его конфигурации.

Этот метод является обратным к get_config, способным восстановить тот же слой из словаря конфигурации. Он не обрабатывает соединение слоев (обрабатывается сетью), а также веса (обрабатываются set_weights).

Аргументы
config Словарь Python, как правило, вывод get_config.
Возвращает
Экземпляр слоя.

symbolic_call

Просмотреть исходный код

symbolic_call(
    *args, **kwargs
)

© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/api_docs/python/tf/keras/layers/FlaxLayer

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API