mirror of
https://github.com/tensorflow/tensorflow.git
synced 2026-09-28 13:23:37 +08:00
Internal cleanup.
PiperOrigin-RevId: 180578376
This commit is contained in:
committed by
TensorFlower Gardener
parent
a071dd520e
commit
138ce57600
@@ -238,7 +238,8 @@ class RNNCell(base_layer.Layer):
|
||||
# Try to use the last cached zero_state. This is done to avoid recreating
|
||||
# zeros, especially when eager execution is enabled.
|
||||
state_size = self.state_size
|
||||
if hasattr(self, "_last_zero_state"):
|
||||
is_eager = context.in_eager_mode()
|
||||
if is_eager and hasattr(self, "_last_zero_state"):
|
||||
(last_state_size, last_batch_size, last_dtype,
|
||||
last_output) = getattr(self, "_last_zero_state")
|
||||
if (last_batch_size == batch_size and
|
||||
@@ -247,7 +248,8 @@ class RNNCell(base_layer.Layer):
|
||||
return last_output
|
||||
with ops.name_scope(type(self).__name__ + "ZeroState", values=[batch_size]):
|
||||
output = _zero_state_tensors(state_size, batch_size, dtype)
|
||||
self._last_zero_state = (state_size, batch_size, dtype, output)
|
||||
if is_eager:
|
||||
self._last_zero_state = (state_size, batch_size, dtype, output)
|
||||
return output
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user