@@ -174,7 +174,7 @@ def first_step(self):
174174
175175
176176 self .J = np .zeros ( self .gridsize )
177- self .action_policy = np .zeros ( self .gridsize )
177+ self .action_policy = np .zeros ( self .gridsize , dtype = int )
178178 self .u0_policy = np .zeros ( self .gridsize )
179179 self .Jnew = np .zeros ( self .gridsize )
180180 self .Jplot = np .zeros ( self .gridsize )
@@ -437,7 +437,7 @@ def load_data(self, name = 'DP_data'):
437437 # Dyan prog data
438438 self .X = np .load ( name + '_X' + '.npy' )
439439 self .J = np .load ( name + '_J' + '.npy' )
440- self .action_policy = np .load ( name + '_a' + '.npy' )
440+ self .action_policy = np .load ( name + '_a' + '.npy' ). astype ( int )
441441 self .u0_policy = np .load ( name + '_u0' + '.npy' )
442442
443443 except :
@@ -450,10 +450,10 @@ def save_data(self, name = 'DP_data'):
450450 """ Save optimal controller policy and cost to go """
451451
452452 # Dyan prog data
453- np .save ( name + '_X' , self .X )
454- np .save ( name + '_J' , self .J )
455- np .save ( name + '_a' , self .action_policy )
456- np .save ( name + '_u0' , self .u0_policy )
453+ np .save ( name + '_X' , self .X )
454+ np .save ( name + '_J' , self .J )
455+ np .save ( name + '_a' , self .action_policy . astype ( int ) )
456+ np .save ( name + '_u0' , self .u0_policy )
457457
458458 ################################
459459 def compute_traj_cost (self ):
0 commit comments