# \[rllib\] SampleBatch "state\_in\_0" dimension shorter than expected

**URL:** <https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354>\
**Category:** RLlib\
**Created:** [May 31, 2021, 8:33pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354 "2021-05-31T20:33:42Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![rw422scarlet](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/rw422scarlet/32/950_2.png) [@rw422scarlet](https://discuss.ray.io/u/rw422scarlet)\
**Post date:** [May 31, 2021, 8:33pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/1 "2021-05-31T20:33:42Z")

</div>

Hi guys,

I am using sort of a Bayesian filter (a VAE to be exact) instead of an RNN but I would like to use the “state” variable in the RNN interface to store the belief state and I would like to directly access the entire history of (prior) belief state for training.

I though the “state\_in\_0” column in SampleBatch is used exactly for this purpose, so I added this to my view requirement as follow:

```
self.view_requirements["state_in_0"] = ViewRequirement(
        data_col="state_out_0",
        shift=-1,
        used_for_training=True,
    )

```

and I initialized a dummy initial state with zeros and replaced it with a belief state in the action sampler function:

```
def get_initial_state(self):
    return torch.zeros(1, dummy_dim)

def action_sampler_fn(policy, model, input_dict, state, timestep):
    if tilmestep == 0:
        state = make_actual_belief_state()
    action, logp = model(input_dict, state)
    return action, logo, state

```

However, when I check sample batch when computing losses, the length of state\_in\_0 is the number of rollout episodes, while the length of state\_out\_0 is the number time steps, which is what I wanted.

Because this is very confusing, I would like to get some clarity on the stored variables, when they are stored, and what they are used for.

Thanks

---

<div class="post-metadata">

**Author:** ![smorad](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/smorad/32/272_2.png) [@smorad](https://discuss.ray.io/u/smorad)\
**Post date:** [June 1, 2021, 12:32pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/2 "2021-06-01T12:32:39Z")

</div>

Your state cannot change shape. It must always be (B,1,dummy\_dim) (the shape from get\_initial\_state)

---

<div class="post-metadata">

**Author:** ![rw422scarlet](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/rw422scarlet/32/950_2.png) [@rw422scarlet](https://discuss.ray.io/u/rw422scarlet)\
**Post date:** [June 1, 2021, 1:27pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/3 "2021-06-01T13:27:22Z")

</div>

My expected shape is (episode\_timesteps, 1 dummy\_dim), but like you said I get (num\_rollout, 1, dummy\_dim).

---

<div class="post-metadata">

**Author:** ![mannyv](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/mannyv/32/606_2.png) [@mannyv](https://discuss.ray.io/u/mannyv)\
**Post date:** [June 1, 2021, 10:36pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/4 "2021-06-01T22:36:26Z")

</div>

What is your max\_seq\_len? Before passing the model into forwad\_rnn two things happen. 1. The data in the sample batch is padded so that all inputs are the same size as your longest episode. So if you had sequences of length [5,12,10,4] they would all be padded with zeros to be 12 steps long. You would end up with a total of12 \* 4 timesteps in your sample batch.

They are also shortened to be no larger than max sequence length. Let’s say your max\_seq\_len in the model config was 20. Your state in would be of size [4,cell\_size] since from the rnn perspective you do not need to truncate backprop through time. This is likely what you are seeing. When you are using an rnn you only get the initial state of the sequence. The other states are generated internally by the rnn logic. If your max sequence length was 5 in the example above your would likely have seq\_lens of [5,5,5,2,5,5,4] , a sample batch that was padded to 7\*5 and a state\_in with a shape of [7,cell\_size]

The other thing to keep in mind is that there are several passes though the models with dummy data before training starts to calculate view requirements and other values needed by compute action and the loss functions. I have found that sometimes those passes have shapes that I never see during actual traing.

---

<div class="post-metadata">

**Author:** ![rw422scarlet](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/rw422scarlet/32/950_2.png) [@rw422scarlet](https://discuss.ray.io/u/rw422scarlet)\
**Post date:** [June 4, 2021, 2:18pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/5 "2021-06-04T14:18:03Z")

</div>

Hi thanks for the clarification. I think I got the hang of it now. So state\_in only stores the initial hidden cell, then hidden cells generated internally by the RNN will be stored in state\_out, is this correct? And then in the case that episodes are of different length, when they appear in your compute\_loss\_fn() they will be padded to equal length.

---

<div class="post-metadata">

**Author:** ![mannyv](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/mannyv/32/606_2.png) [@mannyv](https://discuss.ray.io/u/mannyv)\
**Post date:** [June 4, 2021, 3:16pm UTC](https://discuss.ray.io/t/rllib-samplebatch-state-in-0-dimension-shorter-than-expected/2354/6 "2021-06-04T15:16:26Z")

</div>

@rw422scarlet,  
Yes. Just a few points of clarification.

During the collection of new episodes, when compute\_actions is called. The state\_in will hold the value of the state\_out of the previous call to compute\_actions. If this is the first step of a new episode this will be whatever is returned by the policy in get\_initial\_state. When compute\_actions returns, it will provide logits (not the actual action those are determined by the action distribution that is applied separately to these values) and the resulting state as state\_out. State is chained this way from timestep to timestep throughout an episode.

Initially all of the state\_in and state\_out are saved but when a sample batch is constructed for the loss function, usually, all of the state\_in and state\_out that correspond to timesteps in an episode where t % max\_seq\_len ==0 are saved and the states between them are discarded to save space.

The other thing that happens, as you noted corectly is that all of the sequences are padded to be the same length as the largest sequence in the samplebatch.

In the loss functions for algorithms that support RNN there is a process after the loss is computed for each step that will zero out the padded timesteps so that they do not contribute to the loss or the backward passes to compute the gradients. Here is an example from PPO.

> <https://github.com/ray-project/ray/blob/ebc44c3d76d114e6192d697e3715aa73bc66924d/rllib/agents/ppo/ppo_torch_policy.py#L97-L101>

> <https://github.com/ray-project/ray/blob/ebc44c3d76d114e6192d697e3715aa73bc66924d/rllib/agents/ppo/ppo_torch_policy.py#L49-L67>
