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