Skip to content

Commit fef2465

Browse files
committed
Fixed unidirectional_sequence_lstm
1 parent 9f5ac25 commit fef2465

1 file changed

Lines changed: 111 additions & 10 deletions

File tree

tensorflow/lite/micro/kernels/cmsis_nn/unidirectional_sequence_lstm.cc

Lines changed: 111 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -270,7 +270,7 @@ TfLiteStatus CMSIS_NN_PortOpData(TfLiteContext* context, OpDataLSTM* params_ref,
270270
}
271271

272272
TfLiteStatus CMSIS_NN_EvalInteger8x8_16Lstm(
273-
const OpData& op_data, const LSTMKernelContents& kernel_content,
273+
const OpData& op_data, LSTMKernelContents& kernel_content,
274274
const LSTMBuffers<int16_t>& buffers) {
275275
TFLITE_DCHECK(
276276
kernel_content.GetInternalTensor(tflite::kLstmInputTensor)->dims->size >=
@@ -282,21 +282,74 @@ TfLiteStatus CMSIS_NN_EvalInteger8x8_16Lstm(
282282
kernel_content.GetInternalTensor(tflite::kLstmInputTensor));
283283
int8_t* output =
284284
tflite::micro::GetTensorData<int8_t>(kernel_content.output_tensor);
285+
int8_t* hidden_state =
286+
tflite::micro::GetTensorData<int8_t>(kernel_content.HiddenStateTensor());
287+
int16_t* cell_state =
288+
tflite::micro::GetTensorData<int16_t>(kernel_content.CellStateTensor());
285289

286290
// Create lstm buffer struct
287291
cmsis_nn_lstm_context cmsis_buffers;
288292
cmsis_buffers.temp1 = reinterpret_cast<int16_t*>(buffers.buffer0);
289293
cmsis_buffers.temp2 = reinterpret_cast<int16_t*>(buffers.buffer1);
290-
cmsis_buffers.cell_state = reinterpret_cast<int16_t*>(buffers.buffer2);
291-
292-
arm_lstm_unidirectional_s8(input, output, &op_data.params_cmsis_nn,
293-
&cmsis_buffers);
294+
cmsis_buffers.cell_state = cell_state;
295+
296+
const auto& params = op_data.params_cmsis_nn;
297+
298+
#ifdef CMSIS_NN_STATEFUL_LSTM
299+
cmsis_buffers.hidden_state = hidden_state;
300+
arm_cmsis_nn_status status =
301+
arm_lstm_unidirectional_s8(input, output, &params, &cmsis_buffers);
302+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
303+
#else
304+
if (params.time_major) {
305+
int8_t* step_hidden_in = hidden_state;
306+
for (int t = 0; t < params.time_steps; t++) {
307+
const int8_t* data_in =
308+
input + (t * params.batch_size * params.input_size);
309+
int8_t* hidden_out =
310+
output + (t * params.batch_size * params.hidden_size);
311+
312+
arm_cmsis_nn_status status = arm_nn_lstm_step_s8(
313+
data_in, step_hidden_in, hidden_out, &params, &cmsis_buffers, 1);
314+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
315+
step_hidden_in = hidden_out;
316+
}
317+
if (params.time_steps > 0) {
318+
std::copy_n(step_hidden_in, params.batch_size * params.hidden_size,
319+
hidden_state);
320+
}
321+
} else {
322+
cmsis_nn_lstm_params step_params = params;
323+
step_params.batch_size = 1;
324+
for (int b = 0; b < params.batch_size; b++) {
325+
int8_t* step_hidden_in = hidden_state + b * params.hidden_size;
326+
cmsis_buffers.cell_state = cell_state + b * params.hidden_size;
327+
328+
for (int t = 0; t < params.time_steps; t++) {
329+
const int8_t* data_in =
330+
input + (b * params.time_steps + t) * params.input_size;
331+
int8_t* hidden_out =
332+
output + (b * params.time_steps + t) * params.hidden_size;
333+
334+
arm_cmsis_nn_status status =
335+
arm_nn_lstm_step_s8(data_in, step_hidden_in, hidden_out,
336+
&step_params, &cmsis_buffers, 1);
337+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
338+
step_hidden_in = hidden_out;
339+
}
340+
if (params.time_steps > 0) {
341+
std::copy_n(step_hidden_in, params.hidden_size,
342+
hidden_state + b * params.hidden_size);
343+
}
344+
}
345+
}
346+
#endif
294347

295348
return kTfLiteOk;
296349
}
297350

298351
TfLiteStatus CMSIS_NN_EvalInteger16x8_16Lstm(
299-
const OpData& op_data, const LSTMKernelContents& kernel_content,
352+
const OpData& op_data, LSTMKernelContents& kernel_content,
300353
const LSTMBuffers<int16_t>& buffers) {
301354
TFLITE_DCHECK(
302355
kernel_content.GetInternalTensor(tflite::kLstmInputTensor)->dims->size >=
@@ -308,15 +361,63 @@ TfLiteStatus CMSIS_NN_EvalInteger16x8_16Lstm(
308361
kernel_content.GetInternalTensor(tflite::kLstmInputTensor));
309362
int16_t* output =
310363
tflite::micro::GetTensorData<int16_t>(kernel_content.output_tensor);
364+
int16_t* hidden_state =
365+
tflite::micro::GetTensorData<int16_t>(kernel_content.HiddenStateTensor());
366+
int16_t* cell_state =
367+
tflite::micro::GetTensorData<int16_t>(kernel_content.CellStateTensor());
311368

312369
// Create lstm buffer struct
313370
cmsis_nn_lstm_context cmsis_buffers;
314371
cmsis_buffers.temp1 = reinterpret_cast<int16_t*>(buffers.buffer0);
315372
cmsis_buffers.temp2 = reinterpret_cast<int16_t*>(buffers.buffer1);
316-
cmsis_buffers.cell_state = reinterpret_cast<int16_t*>(buffers.buffer2);
317-
318-
arm_lstm_unidirectional_s16(input, output, &op_data.params_cmsis_nn,
319-
&cmsis_buffers);
373+
cmsis_buffers.cell_state = cell_state;
374+
375+
const auto& params = op_data.params_cmsis_nn;
376+
377+
#ifdef CMSIS_NN_STATEFUL_LSTM
378+
cmsis_buffers.hidden_state = hidden_state;
379+
arm_cmsis_nn_status status =
380+
arm_lstm_unidirectional_s16(input, output, &params, &cmsis_buffers);
381+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
382+
#else
383+
if (params.time_major) {
384+
for (int t = 0; t < params.time_steps; t++) {
385+
const int16_t* data_in =
386+
input + (t * params.batch_size * params.input_size);
387+
int16_t* hidden_out =
388+
output + (t * params.batch_size * params.hidden_size);
389+
390+
arm_cmsis_nn_status status = arm_nn_lstm_step_s16(
391+
data_in, hidden_state, hidden_out, &params, &cmsis_buffers, 1);
392+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
393+
394+
// Update hidden state for next step
395+
std::copy_n(hidden_out, params.batch_size * params.hidden_size,
396+
hidden_state);
397+
}
398+
} else {
399+
cmsis_nn_lstm_params step_params = params;
400+
step_params.batch_size = 1;
401+
for (int b = 0; b < params.batch_size; b++) {
402+
for (int t = 0; t < params.time_steps; t++) {
403+
const int16_t* data_in =
404+
input + (b * params.time_steps + t) * params.input_size;
405+
int16_t* hidden_out =
406+
output + (b * params.time_steps + t) * params.hidden_size;
407+
int16_t* current_hidden = hidden_state + b * params.hidden_size;
408+
cmsis_buffers.cell_state = cell_state + b * params.hidden_size;
409+
410+
arm_cmsis_nn_status status =
411+
arm_nn_lstm_step_s16(data_in, current_hidden, hidden_out,
412+
&step_params, &cmsis_buffers, 1);
413+
if (status != ARM_CMSIS_NN_SUCCESS) return kTfLiteError;
414+
415+
// Update hidden state for next step
416+
std::copy_n(hidden_out, params.hidden_size, current_hidden);
417+
}
418+
}
419+
}
420+
#endif
320421

321422
return kTfLiteOk;
322423
}

0 commit comments

Comments
 (0)