tf.contrib.training.SequenceQueueingStateSaver
SequenceQueueingStateSaver предоставляет доступ к состоянию значений из входных данных.
tf.contrib.training.SequenceQueueingStateSaver(
batch_size, num_unroll, input_length, input_key, input_sequences, input_context,
initial_states, capacity=None, allow_small_batch=False, name=None
)
Этот класс предназначен для использования вместо, например, Queue, для разделения входных данных последовательностей переменной длины на сегменты последовательностей фиксированной длины и их объединения в мини-пакеты. Он сохраняет контексты и состояние для последовательности в сегментах. Его можно использовать совместно с QueueRunner (см. пример ниже).
SequenceQueueingStateSaver (SQSS) принимает по одному примеру за раз через входы input_length, input_key, input_sequences (словарь), input_context (словарь) и initial_states (словарь). Последовательности, значения в input_sequences, могут иметь переменную первую размерность (padded_length), хотя эта размерность всегда должна быть кратной num_unroll. Все остальные размерности должны быть фиксированными и доступными через вызовы get_shape. Длина до заполнения может быть записана в input_length. Значения контекста в input_context должны иметь фиксированные и хорошо определенные размерности. Начальные значения состояния должны иметь фиксированные и хорошо определенные размерности.
SQSS разбивает последовательности входного примера на сегменты длиной num_unroll. Между примерами формируются мини-пакеты размером batch_size. Эти мини-пакеты содержат сегмент последовательностей, копируют значения контекста и сохраняют информацию о состоянии, длине и ключе исходных входных примеров. В первом сегменте примера состояние всё ещё является начальным состоянием. Его можно обновить; обновленные значения состояния доступны в последующих сегментах того же примера. После каждого сегмента необходимо вызвать batch.save_state(), что выполняется state_saving_rnn. Без этого вызова операция декьюинга, связанная с SQSS, не будет выполнена. Внутренне SQSS имеет очередь для входных примеров. Её capacity настраивается. Если он меньше batch_size, операция декьюинга будет блокироваться неопределённо долго. Небольшое кратное значение batch_size - хороший ориентир для предотвращения превращения очереди в узкое место и замедления обучения. Если он слишком большой (по умолчанию он не ограничен), увеличивается потребление памяти. Кроме того, при многократном повторении одних и тех же входных примеров с использованием одного и того же key значение capacity должно быть меньше количества примеров.
Префечер, который считывает одну развернутую последовательность входных данных переменной длины за раз, доступен через prefetch_op. Базовый объект Barrier доступен через barrier. Обработанные мини-пакеты, а также возможности чтения и записи состояния доступны через next_batch. В частности, next_batch предоставляет доступ ко всем данным мини-пакетов, включая следующее, см. NextQueuedSequenceBatch для подробностей:
-
total_length,length,insertion_index,key,next_key, -
sequence(индекс индекса сегмента времени каждого элемента мини-пакета), -
sequence_count(общее количество сегментов времени для каждого элемента мини-пакета), -
context(словарь скопированных значений контекста мини-пакета), -
sequences(словарь разделенных мини-пакетированных последовательностей переменной длины), -
state(для доступа к состояниям текущих сегментов этих элементов) -
save_state(для сохранения состояний следующих сегментов этих элементов)
Пример использования:
batch_size = 32
num_unroll = 20
lstm_size = 8
cell = tf.compat.v1.nn.rnn_cell.BasicLSTMCell(num_units=lstm_size)
initial_state_values = tf.zeros(cell.state_size, dtype=tf.float32)
raw_data = get_single_input_from_input_reader()
length, key, sequences, context = my_parser(raw_data)
assert "input" in sequences.keys()
assert "label" in context.keys()
initial_states = {"lstm_state": initial_state_value}
stateful_reader = tf.SequenceQueueingStateSaver(
batch_size, num_unroll,
length=length, input_key=key, input_sequences=sequences,
input_context=context, initial_states=initial_states,
capacity=batch_size*100)
batch = stateful_reader.next_batch
inputs = batch.sequences["input"]
context_label = batch.context["label"]
inputs_by_time = tf.split(value=inputs, num_or_size_splits=num_unroll, axis=1)
assert len(inputs_by_time) == num_unroll
lstm_output, _ = tf.contrib.rnn.static_state_saving_rnn(
cell,
inputs_by_time,
state_saver=batch,
state_name="lstm_state")
# Start a prefetcher in the background
sess = tf.compat.v1.Session()
num_threads = 3
queue_runner = tf.compat.v1.train.QueueRunner(
stateful_reader, [stateful_reader.prefetch_op] * num_threads)
tf.compat.v1.train.add_queue_runner(queue_runner)
tf.compat.v1.train.start_queue_runners(sess=session)
while True:
# Step through batches, perform training or inference...
session.run([lstm_output])
Примечание: Обычно барьер предоставляется QueueRunner, как в примерах выше. QueueRunner закроет барьер, если префеч-операция получит ошибку OutOfRange от входных очередей выше по потоку (т.е., достигнет конца входных данных). Если барьер закрыт, новые примеры больше не добавляются в SQSS. Однако базовый барьер может всё ещё содержать дальнейшие шаги развертки примеров, которые не прошли все итерации. Для корректного завершения всех примеров флаг allow_small_batch должен быть установлен в true, что заставит SQSS выдавать постепенно уменьшающиеся мини-пакеты с оставшимися примерами.
| Аргументы | |
|---|---|
batch_size | целочисленное или 32-битовое скалярное значение Tensor, размер мини-пакетов при обращении к методу state() и свойствам context, sequences и т. д. |
num_unroll | целое число Python, количество временных шагов для развертки за раз. Последовательности входных данных длиной k затем разбиваются на k / num_unroll сегментов. |
input_length | 32-битовое скалярное значение Tensor, длина последовательности до заполнения. Это значение может быть не более padded_length для любого данного входного значения (см. определение padded_length ниже). Сгруппированные и общие длины текущей итерации доступны через свойства length и total_length. Форма входной длины (скаляр) должна быть полностью указана. |
input_key | строковое скалярное значение Tensor, уникальный ключ для данного входного значения. Он используется для отслеживания разделенных элементов мини-пакета этого входного значения. Сгруппированные ключи текущей итерации доступны через свойство key. Форма input_key (скаляр) должна быть полностью указана. |
input_sequences | словарь, сопоставляющий строковые имена с значениями Tensor. Значения должны иметь совпадающую первую размерность, называемую padded_length. SequenceQueueingStateSaver разделит эти тензоры по этой первой размерности на элементы мини-пакета с размерностью num_unroll. Сгруппированные и сегментированные последовательности текущей итерации доступны через свойство sequences. Примечание: |
input_context | словарь, сопоставляющий строковые имена со значениями Tensor. Значения обрабатываются как «глобальные» для всех временных разделов данного входного значения и будут копироваться для всех элементов мини-пакета соответственно. Сгруппированный и скопированный контекст текущей итерации доступен через свойство context.Примечание: Все входные значения input_context должны иметь полностью определенные формы. |
initial_states | словарь, сопоставляющий строковые имена состояний с многомерными значениями (например, константы или тензоры). Этот вход определяет набор состояний, которые будут отслеживаться во время вычисления итераций, и к которым можно получить доступ через методы state и save_state.Примечание: Все начальные значения initial_state должны иметь полностью определенные формы. |
capacity | максимальная ёмкость SQSS в количестве примеров. Должно быть как минимум batch_size. По умолчанию не ограничено. |
allow_small_batch | Если значение true, SQSS возвращает меньшие пакеты, когда нет достаточного количества входных примеров для заполнения целого пакета, и достигнут конец входных данных (т.е., базовый барьер закрыт). |
name | строка имени операции (необязательно). |
| Возбуждения | |
|---|---|
TypeError | если любой из входов не является ожидаемого типа. |
ValueError | если любое из входных значений несовместимо, например, если из входов недостаточно информации о форме для построения сохранителя состояния. |
| Атрибуты | |
|---|---|
barrier | |
batch_size | |
name | |
next_batch | NextQueuedSequenceBatch предоставляющий доступ к данным сгруппированного вывода. Также предоставляет доступ к методам Для доступа к данным в |
num_unroll | |
prefetch_op | Операция, используемая для предварительной загрузки новых данных в сохранитель состояния. Выполнение его один раз помещает один новый входной пример в сохранитель состояния. В первый раз при вызове он дополнительно создает префеч-операцию. Последующие вызовы просто возвращают ранее созданный Его следует запускать в отдельном потоке, например, с помощью |
Методы
close
close(
cancel_pending_enqueues=False, name=None
)
Закрывает барьер и FIFOQueue.
Эта операция сигнализирует о том, что больше сегменты новых последовательностей не будут помещены в очередь. Новые сегменты уже вставленных последовательностей всё ещё могут быть помещены в очередь и декьюированы, если есть достаточное количество для заполнения пакета или если allow_small_batch имеет значение true. В противном случае операции декьюинга завершатся немедленно.
| Аргументы | |
|---|---|
cancel_pending_enqueues | (Необязательно.) Булево значение, по умолчанию False. Если True, все ожидающие помещения в очередь в подчинённые очереди будут отменены, и завершение уже начатых последовательностей невозможно. |
name | Необязательное имя для операции. |
| Возвращает | |
|---|---|
| Операция, закрывающая барьер и FIFOQueue. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/contrib/training/SequenceQueueingStateSaver