@@ -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