Skip to content

Commit 0386c49

Browse files
mbzCopybara-Service
authored andcommitted
Separating the reward model.
PiperOrigin-RevId: 204533573
1 parent 0416dfc commit 0386c49

1 file changed

Lines changed: 116 additions & 59 deletions

File tree

tensor2tensor/models/research/next_frame.py

Lines changed: 116 additions & 59 deletions
Original file line numberDiff line numberDiff line change
@@ -221,18 +221,102 @@ def construct_latent_tower(self, images):
221221

222222
return mean, std
223223

224-
def reward_prediction(self, inputs):
224+
def bottom_part_tower(self, input_image, input_reward, action, latent,
225+
lstm_state, lstm_size, conv_size):
226+
"""The bottom part of predictive towers.
227+
228+
With the current (early) design, the main prediction tower and
229+
the reward prediction tower share the same arcitecture. TF Scope can be
230+
adjusted as required to either share or not share the weights between
231+
the two towers.
232+
233+
Args:
234+
input_image: the current image.
235+
input_reward: the current reward.
236+
action: the action taken by the agent.
237+
latent: the latent vector.
238+
lstm_state: the current internal states of conv lstms.
239+
lstm_size: the size of lstms.
240+
conv_size: the size of convolutions.
241+
242+
Returns:
243+
- the output of the partial network.
244+
- intermidate outputs for skip connections.
245+
"""
246+
layer_norm = tf.contrib.layers.layer_norm
247+
lstm_func = self.conv_lstm_2d
248+
249+
input_image = common_layers.make_even_size(input_image)
250+
enc0 = slim.layers.conv2d(
251+
input_image,
252+
conv_size[0], [5, 5],
253+
stride=2,
254+
scope="scale1_conv1",
255+
normalizer_fn=layer_norm,
256+
normalizer_params={"scope": "layer_norm1"})
257+
258+
hidden1, lstm_state[0] = lstm_func(
259+
enc0, lstm_state[0], lstm_size[0], scope="state1")
260+
hidden1 = layer_norm(hidden1, scope="layer_norm2")
261+
hidden2, lstm_state[1] = lstm_func(
262+
hidden1, lstm_state[1], lstm_size[1], scope="state2")
263+
hidden2 = layer_norm(hidden2, scope="layer_norm3")
264+
hidden2 = common_layers.make_even_size(hidden2)
265+
enc1 = slim.layers.conv2d(
266+
hidden2, hidden2.get_shape()[3], [3, 3], stride=2, scope="conv2")
267+
268+
hidden3, lstm_state[2] = lstm_func(
269+
enc1, lstm_state[2], lstm_size[2], scope="state3")
270+
hidden3 = layer_norm(hidden3, scope="layer_norm4")
271+
hidden4, lstm_state[3] = lstm_func(
272+
hidden3, lstm_state[3], lstm_size[3], scope="state4")
273+
hidden4 = layer_norm(hidden4, scope="layer_norm5")
274+
hidden4 = common_layers.make_even_size(hidden4)
275+
enc2 = slim.layers.conv2d(
276+
hidden4, hidden4.get_shape()[3], [3, 3], stride=2, scope="conv3")
277+
278+
# Pass in reward and action.
279+
emb_action = self.encode_to_shape(action, enc2.get_shape(), "action_enc")
280+
emb_reward = self.encode_to_shape(
281+
input_reward, enc2.get_shape(), "reward_enc")
282+
enc2 = tf.concat(axis=3, values=[enc2, emb_action, emb_reward])
283+
284+
if latent is not None:
285+
with tf.control_dependencies([latent]):
286+
enc2 = tf.concat([enc2, latent], 3)
287+
288+
enc3 = slim.layers.conv2d(
289+
enc2, hidden4.get_shape()[3], [1, 1], stride=1, scope="conv4")
290+
291+
hidden5, lstm_state[4] = lstm_func(
292+
enc3, lstm_state[4], lstm_size[4], scope="state5") # last 8x8
293+
hidden5 = layer_norm(hidden5, scope="layer_norm6")
294+
295+
return hidden5, (enc0, enc1)
296+
297+
def reward_prediction(
298+
self, input_image, input_reward, action, lstm_state, latent):
225299
"""Builds a reward prediction network."""
226-
conv_size = self.tinyify([32, 16, 1])
300+
conv_size = self.tinyify([32, 32, 16, 4])
301+
lstm_size = self.tinyify([32, 64, 128, 64, 32])
302+
227303
with tf.variable_scope("reward_pred", reuse=tf.AUTO_REUSE):
228-
x = inputs
304+
hidden5, _ = self.bottom_part_tower(
305+
input_image, input_reward, action, latent,
306+
lstm_state, lstm_size, conv_size)
307+
308+
x = hidden5
229309
x = slim.batch_norm(x, scope="reward_bn0")
230-
x = slim.conv2d(x, conv_size[0], [3, 3], scope="reward_conv1")
310+
x = slim.conv2d(x, conv_size[1], [3, 3], scope="reward_conv1")
231311
x = slim.batch_norm(x, scope="reward_bn1")
232-
x = slim.conv2d(x, conv_size[1], [3, 3], scope="reward_conv2")
312+
x = slim.conv2d(x, conv_size[2], [3, 3], scope="reward_conv2")
233313
x = slim.batch_norm(x, scope="reward_bn2")
234-
x = slim.conv2d(x, conv_size[2], [3, 3], scope="reward_conv3")
235-
return x
314+
x = slim.conv2d(x, conv_size[3], [3, 3], scope="reward_conv3")
315+
316+
pred_reward = self.decode_to_shape(
317+
x, input_reward.shape, "reward_dec")
318+
319+
return pred_reward, lstm_state
236320

237321
def encode_to_shape(self, inputs, shape, scope):
238322
"""Encode the given tensor to given image shape."""
@@ -280,51 +364,11 @@ def construct_predictive_tower(
280364
img_height, img_width, color_channels = self.hparams.problem.frame_shape
281365

282366
with tf.variable_scope("main", reuse=tf.AUTO_REUSE):
283-
input_image = common_layers.make_even_size(input_image)
284-
enc0 = slim.layers.conv2d(
285-
input_image,
286-
conv_size[0], [5, 5],
287-
stride=2,
288-
scope="scale1_conv1",
289-
normalizer_fn=layer_norm,
290-
normalizer_params={"scope": "layer_norm1"})
291-
292-
hidden1, lstm_state[0] = lstm_func(
293-
enc0, lstm_state[0], lstm_size[0], scope="state1")
294-
hidden1 = layer_norm(hidden1, scope="layer_norm2")
295-
hidden2, lstm_state[1] = lstm_func(
296-
hidden1, lstm_state[1], lstm_size[1], scope="state2")
297-
hidden2 = layer_norm(hidden2, scope="layer_norm3")
298-
hidden2 = common_layers.make_even_size(hidden2)
299-
enc1 = slim.layers.conv2d(
300-
hidden2, hidden2.get_shape()[3], [3, 3], stride=2, scope="conv2")
301-
302-
hidden3, lstm_state[2] = lstm_func(
303-
enc1, lstm_state[2], lstm_size[2], scope="state3")
304-
hidden3 = layer_norm(hidden3, scope="layer_norm4")
305-
hidden4, lstm_state[3] = lstm_func(
306-
hidden3, lstm_state[3], lstm_size[3], scope="state4")
307-
hidden4 = layer_norm(hidden4, scope="layer_norm5")
308-
hidden4 = common_layers.make_even_size(hidden4)
309-
enc2 = slim.layers.conv2d(
310-
hidden4, hidden4.get_shape()[3], [3, 3], stride=2, scope="conv3")
311-
312-
# Pass in reward and action.
313-
emb_action = self.encode_to_shape(action, enc2.get_shape(), "action_enc")
314-
emb_reward = self.encode_to_shape(
315-
input_reward, enc2.get_shape(), "reward_enc")
316-
enc2 = tf.concat(axis=3, values=[enc2, emb_action, emb_reward])
317-
318-
if latent is not None:
319-
with tf.control_dependencies([latent]):
320-
enc2 = tf.concat([enc2, latent], 3)
321-
322-
enc3 = slim.layers.conv2d(
323-
enc2, hidden4.get_shape()[3], [1, 1], stride=1, scope="conv4")
324-
325-
hidden5, lstm_state[4] = lstm_func(
326-
enc3, lstm_state[4], lstm_size[4], scope="state5") # last 8x8
327-
hidden5 = layer_norm(hidden5, scope="layer_norm6")
367+
hidden5, skips = self.bottom_part_tower(
368+
input_image, input_reward, action, latent,
369+
lstm_state, lstm_size, conv_size)
370+
enc0, enc1 = skips
371+
328372
enc4 = slim.layers.conv2d_transpose(
329373
hidden5, hidden5.get_shape()[3], 3, stride=2, scope="convt1")
330374

@@ -404,11 +448,7 @@ def construct_predictive_tower(
404448
for layer, mask in zip(transformed, mask_list[1:]):
405449
output += layer * mask
406450

407-
p_reward = self.reward_prediction(hidden5)
408-
p_reward = self.decode_to_shape(
409-
p_reward, input_reward.shape, "reward_dec")
410-
411-
return output, p_reward, lstm_state
451+
return output, lstm_state
412452

413453
def get_guassian_latent(self, latent_mean, latent_std):
414454
latent = tf.random_normal(tf.shape(latent_mean), 0, 1, dtype=tf.float32)
@@ -445,6 +485,7 @@ def construct_model(self,
445485

446486
# LSTM states.
447487
lstm_state = [None] * 7
488+
reward_lstm_state = [None] * 5
448489

449490
# Latent tower
450491
if self.hparams.stochastic_model:
@@ -466,8 +507,15 @@ def construct_model(self,
466507
latent = self.get_guassian_latent(latent_mean, latent_std)
467508

468509
# Prediction
469-
pred_image, pred_reward, lstm_state = self.construct_predictive_tower(
510+
pred_image, lstm_state = self.construct_predictive_tower(
470511
input_image, input_reward, action, lstm_state, latent)
512+
513+
if self.hparams.reward_prediction:
514+
pred_reward, reward_lstm_state = self.reward_prediction(
515+
input_image, input_reward, action, reward_lstm_state, latent)
516+
else:
517+
pred_reward = input_reward
518+
471519
gen_images.append(pred_image)
472520
gen_rewards.append(pred_reward)
473521

@@ -733,6 +781,7 @@ def construct_model(self, images, actions, rewards):
733781

734782
# LSTM states.
735783
lstm_state = [None] * 7
784+
reward_lstm_state = [None] * 5
736785

737786
pred_image, pred_reward, latent = None, None, None
738787
for timestep, image, action, reward in zip(
@@ -753,8 +802,15 @@ def construct_model(self, images, actions, rewards):
753802
latent_stds.append(latent_std)
754803

755804
# Prediction
756-
pred_image, pred_reward, lstm_state = self.construct_predictive_tower(
805+
pred_image, lstm_state = self.construct_predictive_tower(
757806
input_image, input_reward, action, lstm_state, latent)
807+
808+
if self.hparams.reward_prediction:
809+
pred_reward, reward_lstm_state = self.reward_prediction(
810+
input_image, input_reward, action, reward_lstm_state, latent)
811+
else:
812+
pred_reward = input_reward
813+
758814
gen_images.append(pred_image)
759815
gen_rewards.append(pred_reward)
760816

@@ -1064,6 +1120,7 @@ def next_frame_stochastic():
10641120
hparams.input_modalities = "inputs:video:l2raw"
10651121
hparams.video_modality_loss_cutoff = 0.0
10661122
hparams.add_hparam("stochastic_model", True)
1123+
hparams.add_hparam("reward_prediction", True)
10671124
hparams.add_hparam("model_options", "CDNA")
10681125
hparams.add_hparam("num_masks", 10)
10691126
hparams.add_hparam("latent_channels", 1)

0 commit comments

Comments
 (0)