tf.keras.layers.JaxLayer
Слой Keras, который оборачивает модель JAX.
Наследуется от: Layer, Operation
tf.keras.layers.JaxLayer(
call_fn, init_fn=None, params=None, state=None, seed=None, **kwargs
)
Этот слой позволяет использовать компоненты JAX в Keras, когда JAX используется в качестве бэкенда для Keras.
Функция модели
Этот слой принимает модели JAX в виде функции, call_fn, которая должна принимать следующие аргументы с этими точными именами:
-
params: обучаемые параметры модели. -
state(необязательно): необучаемое состояние модели. Можно опустить, если у модели нет не обучаемого состояния. -
rng(необязательно): экземплярjax.random.PRNGKey. Можно опустить, если модели не нужны генераторы случайных чисел ни во время обучения, ни во время вывода. -
inputs: входные данные в модель, массив JAX илиPyTreeмассивов. -
training(необязательно): аргумент, указывающий, находимся ли мы в режиме обучения или вывода,Trueпередается в режиме обучения. Можно опустить, если модель ведет себя одинаково в режиме обучения и вывода.
Аргумент inputs является обязательным. Входные данные модели должны быть предоставлены через один аргумент. Если модель JAX принимает несколько входных данных в виде отдельных аргументов, их необходимо объединить в одну структуру, например, в tuple или в dict.
Инициализация весов модели
Инициализация params и state модели может обрабатываться этим слоем, в этом случае должен быть предоставлен аргумент init_fn. Это позволяет модели динамически инициализироваться с правильной формой. В качестве альтернативы, и если форма известна, можно использовать аргумент params и необязательно аргумент state, чтобы создать уже инициализированную модель.
Функция init_fn, если она предоставлена, должна принимать следующие аргументы с этими точными именами:
-
rng: экземплярjax.random.PRNGKey. -
inputs: массив JAX илиPyTreeмассивов со значениями-заполнителями, чтобы предоставить форму входных данных. -
training(необязательно): аргумент, указывающий, находимся ли мы в режиме обучения или вывода.Trueвсегда передается вinit_fn. Можно опустить, независимо от того, есть ли уcall_fnаргументtraining.
Модели с не обучаемым состоянием
Для моделей JAX, у которых есть не обучаемое состояние:
-
call_fnдолжен иметь аргументstate -
call_fnдолжен возвращатьtuple, содержащий выходные данные модели и новое не обучаемое состояние модели -
init_fnдолжен возвращатьtuple, содержащий начальные обучаемые параметры модели и начальное не обучаемое состояние модели.
Этот код показывает возможную комбинацию сигнатур call_fn и init_fn для модели с не обучаемым состоянием. В этом примере модель имеет аргумент training и аргумент rng в call_fn.
def stateful_call(params, state, rng, inputs, training):
outputs = ...
new_state = ...
return outputs, new_state
def stateful_init(rng, inputs):
initial_params = ...
initial_state = ...
return initial_params, initial_state
Модели без не обучаемого состояния
Для моделей JAX без не обучаемого состояния:
-
call_fnне должен иметь аргументstate -
call_fnдолжен возвращать только выходные данные модели -
init_fnдолжен возвращать только начальные обучаемые параметры модели.
Этот код показывает возможную комбинацию сигнатур call_fn и init_fn для модели без не обучаемого состояния. В этом примере модель не имеет аргумента training и не имеет аргумента rng в call_fn.
def stateless_call(params, inputs):
outputs = ...
return outputs
def stateless_init(rng, inputs):
initial_params = ...
return initial_params
Соответствие требуемой сигнатуре
Если у модели есть другая сигнатура, чем требуется JaxLayer, можно легко написать метод-обёртку для адаптации аргументов. Этот пример демонстрирует модель, которая имеет несколько входных данных в виде отдельных аргументов, ожидает несколько генераторов случайных чисел в dict и имеет аргумент deterministic с обратным значением training. Для соответствия входные данные объединяются в одну структуру с помощью tuple, генератор случайных чисел разделяется и используется для заполнения ожидаемого dict, а булево значение инвертируется:
def my_model_fn(params, rngs, input1, input2, deterministic):
...
if not deterministic:
dropout_rng = rngs["dropout"]
keep = jax.random.bernoulli(dropout_rng, dropout_rate, x.shape)
x = jax.numpy.where(keep, x / dropout_rate, 0)
...
...
return outputs
def my_model_wrapper_fn(params, rng, inputs, training):
input1, input2 = inputs
rng1, rng2 = jax.random.split(rng)
rngs = {"dropout": rng1, "preprocessing": rng2}
deterministic = not training
return my_model_fn(params, rngs, input1, input2, deterministic)
keras_layer = JaxLayer(my_model_wrapper_fn, params=initial_params)
Использование с модулями Haiku
JaxLayer позволяет использовать компоненты Haiku в форме haiku.Module. Это достигается путем преобразования модуля в соответствии с шаблоном Haiku и затем передачи module.apply в параметр call_fn и module.init в параметр init_fn, если необходимо.
Если у модели есть не обучаемое состояние, она должна быть преобразована с помощью haiku.transform_with_state. Если у модели нет не обучаемого состояния, она должна быть преобразована с помощью haiku.transform. Кроме того, и необязательно, если модуль не использует генераторы случайных чисел в "apply", он может быть преобразован с помощью haiku.without_apply_rng.
Следующий пример показывает, как создать JaxLayer из модуля Haiku, который использует генераторы случайных чисел через hk.next_rng_key() и принимает аргумент обучения:
class MyHaikuModule(hk.Module):
def __call__(self, x, training):
x = hk.Conv2D(32, (3, 3))(x)
x = jax.nn.relu(x)
x = hk.AvgPool((1, 2, 2, 1), (1, 2, 2, 1), "VALID")(x)
x = hk.Flatten()(x)
x = hk.Linear(200)(x)
if training:
x = hk.dropout(rng=hk.next_rng_key(), rate=0.3, x=x)
x = jax.nn.relu(x)
x = hk.Linear(10)(x)
x = jax.nn.softmax(x)
return x
def my_haiku_module_fn(inputs, training):
module = MyHaikuModule()
return module(inputs, training)
transformed_module = hk.transform(my_haiku_module_fn)
keras_layer = JaxLayer(
call_fn=transformed_module.apply,
init_fn=transformed_module.init,
)
| Args | |
|---|---|
call_fn: Функция для вызова модели. См. описание выше для списка аргументов, которые она принимает, и выходных данных, которые она возвращает. init_fn: функция для вызова для инициализации модели. См. описание выше для списка аргументов, которые она принимает, и выходных данных, которые она возвращает. Если None, то params и/или state должны быть предоставлены. | |
params | A PyTree содержащий все обучаемые параметры модели. Это позволяет передавать обученные параметры или управлять инициализацией. Если и params и state являются None, то init_fn вызывается во время построения для инициализации обучаемых параметров модели. |
state | A PyTree содержащий все не обучаемые состояния модели. Это позволяет передавать изученные состояния или управлять инициализацией. Если и params и state являются None, а call_fn принимает аргумент state, то init_fn вызывается во время построения для инициализации не обучаемого состояния модели. |
seed | Семя генератора случайных чисел. Необязательно. |
| Attributes | |
|---|---|
input | Извлекает тензор(ы) ввода символической операции. Возвращает только тензор(ы), соответствующий первому вызову операции. |
output | Извлекает тензор(ы) вывода слоя. Возвращает только тензор(ы), соответствующий первому вызову операции. |
Методы
from_config
@classmethod
from_config(
config
)
Создает слой из его конфигурации.
Этот метод является обратным get_config, способным восстановить тот же слой из словаря конфигурации. Он не обрабатывает соединение слоев (обрабатывается сетью), а также веса (обрабатываются set_weights).
| Args | |
|---|---|
config | Словарь Python, обычно вывод get_config. |
| Returns | |
|---|---|
| Экземпляр слоя. |
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/JaxLayer