Setting up batches in tensorflow probability sts

Viewed 182

I'm having some trouble setting up fitting and forecasting with tfp.sts. I'm trying to fit the model to be able to take a sequence of 24 observed timesteps and forecast the following 24 steps. For this, I have taken my large sequence of 365 * 24 values and sliced it into 48-step-long batches (so the observed data matrix has shape (364, 48)). I also have another variable for which I am attempting to find linear dependence using tfp.sts.LinearRegression, which has the design matrix with shape (365, 24), or (364, 48, 1) when converted to batches equivalent to the observed data.

I use variational inference to get parameter samples, and here's where I run into a problem - the samples have shapes (n_samples, 364). I then attempt to perform a forecast using these samples. With that, observed data and design matrix, for forecasts have shapes (24, ) and (48, 1), respectively (so as if they were one of the batches when training) and I get a memory allocation error because the forecasting function can't allocate space for a tensor with shape [n_samples,364,24,24]. This makes me think that I've made a mistake in the batch setup, but I can't see where. I attempted, when creating a surrogate posterior distribution to pass the batch_size argument equal to (364, ), but this just crashed as well because of a memory allocation error, it attempted to allocate a very large matrix.

Here's some relevant code...

Model function:

def build_model(observed_time_series, features_design):
    hour_of_day_effect = sts.Seasonal(
        num_seasons=24,
        observed_time_series=observed_time_series,
        name='hour_of_day_effect')

    features_effect = sts.LinearRegression(
        design_matrix=features_design,
        name='temperature_effect')

    autoregressive = sts.Autoregressive(
        order=1,
        observed_time_series=observed_time_series,
        name='autoregressive')

    model = sts.Sum([hour_of_day_effect,
                     # day_of_week_effect,
                     # month_of_year_effect,
                     features_effect,
                     autoregressive,
                     # trend
                    ],
                     observed_time_series=observed_time_series)
    return model

Training function:

def train_vi(model, observed,
             num_variational_steps=200,
             optimizer=tf.optimizers.Adam(learning_rate=.1),
             nsamples=50):

    variational_posteriors = tfp.sts.build_factored_surrogate_posterior(
        model=model)

    @tf.function(experimental_compile=True)
    def train():
        return tfp.vi.fit_surrogate_posterior(
            target_log_prob_fn=model.joint_log_prob(
                observed_time_series=observed),
            surrogate_posterior=variational_posteriors,
            optimizer=optimizer,
            num_steps=num_variational_steps)
    
    global loss_curve
    loss_curve = train()
    
    return variational_posteriors.sample(nsamples)

Forecast function:

def forecast(model, q_samples, observed, steps=24):
    forecast_dist = tfp.sts.forecast(
        model=model,
        observed_time_series=observed,
        parameter_samples=q_samples,
        num_steps_forecast=steps)
    return (forecast_dist.mean().numpy()[..., 0],
            forecast_dist.stddev().numpy()[..., 0])

Data setup

train_lim = '2014-01-01'
df_train = df[:train_lim]
df_test = df[train_lim:]

index_start = 0

design_train = []
observed_train = []
for index_end in range(48, len(df_train), 24):
    df_batch = df_train[index_start:index_end]

    design_train.append(df_batch.air_temperature.values.reshape(-1, 1))
    observed_train.append(df_batch[goal])

    index_start += 24

design_train = tf.stack(design_train)
observed_train = np.array(observed_train)
design_train.shape

Fitting:

model = build_model(observed_train, design_train)
samples = train_vi(model, observed_train)

Forecasting:

day = 15
steps = 24

index_end = day * 24
index_start = index_end - 24

design = df_test.air_temperature[index_start:index_end + steps].values.reshape(-1, 1)
observed = df_test[goal][index_start:index_end].values

model = build_model(observed, design)
means, std = forecast(model, samples, observed, steps=steps) # this fails

And this is the stack trace I get:

---------------------------------------------------------------------------
ResourceExhaustedError                    Traceback (most recent call last)
<ipython-input-18-cddd6738ba4d> in <module>
     11 
     12 model = build_model(observed, design)
---> 13 means, std = forecast(model, samples, observed, steps=steps)
     14 
     15 plt.plot(range(index_start, index_end + steps), df_test[goal][index_start:index_end + steps], label='actual')

<ipython-input-11-4a8440dedb9e> in forecast(model, q_samples, observed, steps)
      1 def forecast(model, q_samples, observed, steps=24):
----> 2     forecast_dist = tfp.sts.forecast(
      3         model=model,
      4         observed_time_series=observed,
      5         parameter_samples=q_samples,

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/sts/forecast.py in forecast(model, observed_time_series, parameter_samples, num_steps_forecast, include_observation_noise)
    313         num_timesteps=num_observed_steps, param_vals=parameter_samples)
    314     (_, _, _, predictive_means, predictive_covs, _, _
--> 315     ) = observed_data_ssm.forward_filter(observed_time_series, mask=mask)
    316 
    317     # Build a batch of state-space models over the forecast period. Because

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/distributions/linear_gaussian_ssm.py in forward_filter(self, x, mask)
    839 
    840     with self._name_and_control_scope('forward_filter', x, {'mask': mask}):
--> 841       return self._forward_filter(x, mask=mask)
    842 
    843   def _forward_filter(self, x, mask=None):

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/distributions/linear_gaussian_ssm.py in _forward_filter(self, x, mask)
    923         self.get_observation_noise_for_timestep)
    924 
--> 925     filter_states = tf.scan(update_step_fn,
    926                             elems=x if mask is None else (x, mask),
    927                             initializer=initial_state)

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/util/dispatch.py in wrapper(*args, **kwargs)
    199     """Call target, and fall back on dispatchers if there is a TypeError."""
    200     try:
--> 201       return target(*args, **kwargs)
    202     except (TypeError, ValueError):
    203       # Note: convert_to_eager_tensor currently raises a ValueError, not a

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/util/deprecation.py in new_func(*args, **kwargs)
    572                   func.__module__, arg_name, arg_value, 'in a future version'
    573                   if date is None else ('after %s' % date), instructions)
--> 574       return func(*args, **kwargs)
    575 
    576     doc = _add_deprecated_arg_value_notice_to_docstring(

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/functional_ops.py in scan_v2(fn, elems, initializer, parallel_iterations, back_prop, swap_memory, infer_shape, reverse, name)
    804     ```
    805   """
--> 806   return scan(
    807       fn=fn,
    808       elems=elems,

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/util/dispatch.py in wrapper(*args, **kwargs)
    199     """Call target, and fall back on dispatchers if there is a TypeError."""
    200     try:
--> 201       return target(*args, **kwargs)
    202     except (TypeError, ValueError):
    203       # Note: convert_to_eager_tensor currently raises a ValueError, not a

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/functional_ops.py in scan(fn, elems, initializer, parallel_iterations, back_prop, swap_memory, infer_shape, reverse, name)
    662       initial_i = i
    663       condition = lambda i, _1, _2: i < n
--> 664     _, _, r_a = control_flow_ops.while_loop(
    665         condition,
    666         compute, (initial_i, a_flat, accs_ta),

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/control_flow_ops.py in while_loop(cond, body, loop_vars, shape_invariants, parallel_iterations, back_prop, swap_memory, name, maximum_iterations, return_same_structure)
   2733                                               list(loop_vars))
   2734       while cond(*loop_vars):
-> 2735         loop_vars = body(*loop_vars)
   2736         if try_to_pack and not isinstance(loop_vars, (list, _basetuple)):
   2737           packed = True

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/control_flow_ops.py in <lambda>(i, lv)
   2724         cond = lambda i, lv: (  # pylint: disable=g-long-lambda
   2725             math_ops.logical_and(i < maximum_iterations, orig_cond(*lv)))
-> 2726         body = lambda i, lv: (i + 1, orig_body(*lv))
   2727       try_to_pack = False
   2728 

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/functional_ops.py in compute(i, a_flat, tas)
    645       packed_elems = input_pack([elem_ta.read(i) for elem_ta in elems_ta])
    646       packed_a = output_pack(a_flat)
--> 647       a_out = fn(packed_a, packed_elems)
    648       nest.assert_same_structure(elems if initializer is None else initializer,
    649                                  a_out)

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/distributions/linear_gaussian_ssm.py in kalman_filter_step(state, elems_t)
   1596     #  u_{t|t-1} = F_t u_{t-1} + b_t
   1597     #  P_{t|t-1} = F_t P_{t-1} F_t' + Q_t
-> 1598     predicted_mean, predicted_cov = kalman_transition(
   1599         filtered_mean,
   1600         filtered_cov,

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/distributions/linear_gaussian_ssm.py in kalman_transition(filtered_mean, filtered_cov, transition_matrix, transition_noise)
   1744                                    transition_matrix,
   1745                                    transition_noise)
-> 1746   predicted_cov = _propagate_cov(filtered_cov,
   1747                                  transition_matrix,
   1748                                  transition_noise)

~/venv/mmes/lib/python3.8/site-packages/tensorflow_probability/python/distributions/linear_gaussian_ssm.py in _propagate_cov(cov, linop, dist)
   1945   """Propagate covariance through linear Gaussian transformation."""
   1946   # For linop A and input cov P, returns `A P A' + dist.cov()`
-> 1947   return linop.matmul(linop.matmul(cov), adjoint_arg=True) + dist.covariance()

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/linalg/linear_operator_block_diag.py in matmul(self, x, adjoint, adjoint_arg, name)
    341                         else self.domain_dimension)
    342         op_dimension.assert_is_compatible_with(x.shape[arg_dim])
--> 343       return self._matmul(x, adjoint=adjoint, adjoint_arg=adjoint_arg)
    344 
    345   def _matmul(self, x, adjoint=False, adjoint_arg=False):

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/linalg/linear_operator_block_diag.py in _matmul(self, x, adjoint, adjoint_arg)
    369     result_list = linear_operator_util.broadcast_matrix_batch_dims(
    370         result_list)
--> 371     return array_ops.concat(result_list, axis=-2)
    372 
    373   def matvec(self, x, adjoint=False, name="matvec"):

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/util/dispatch.py in wrapper(*args, **kwargs)
    199     """Call target, and fall back on dispatchers if there is a TypeError."""
    200     try:
--> 201       return target(*args, **kwargs)
    202     except (TypeError, ValueError):
    203       # Note: convert_to_eager_tensor currently raises a ValueError, not a

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/array_ops.py in concat(values, axis, name)
   1652           dtype=dtypes.int32).get_shape().assert_has_rank(0)
   1653       return identity(values[0], name=name)
-> 1654   return gen_array_ops.concat_v2(values=values, axis=axis, name=name)
   1655 
   1656 

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/ops/gen_array_ops.py in concat_v2(values, axis, name)
   1205       return _result
   1206     except _core._NotOkStatusException as e:
-> 1207       _ops.raise_from_not_ok_status(e, name)
   1208     except _core._FallbackException:
   1209       pass

~/venv/mmes/lib/python3.8/site-packages/tensorflow/python/framework/ops.py in raise_from_not_ok_status(e, name)
   6841   message = e.message + (" name: " + name if name is not None else "")
   6842   # pylint: disable=protected-access
-> 6843   six.raise_from(core._status_to_exception(e.code, message), None)
   6844   # pylint: enable=protected-access
   6845 

~/venv/mmes/lib/python3.8/site-packages/six.py in raise_from(value, from_value)

ResourceExhaustedError: OOM when allocating tensor with shape[50,364,24,24] and type double on /job:localhost/replica:0/task:0/device:GPU:0 by allocator GPU_0_bfc [Op:ConcatV2] name: concat
0 Answers
Related