Spec-Zone.ru › TensorFlow

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

Spec-Zone.ru

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