# Bad inference after perfect training. What am I missing?

**URL:** <https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238>\
**Category:** RLlib\
**Created:** [May 24, 2022, 6:49am UTC](https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238 "2022-05-24T06:49:56Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![vlainic](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/vlainic/32/2330_2.png) [@vlainic](https://discuss.ray.io/u/vlainic)\
**Post date:** [May 24, 2022, 6:49am UTC](https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238/1 "2022-05-24T06:49:56Z")

</div>

Hello,  
so I made a mock-up problem of the real-world combinatorial optimization I have in front of me so I can check the result manually and share it publically.

The `.html` file of the notebook is on [my github](https://htmlpreview.github.io/?https://github.com/vlainic/RLlib-issues/blob/main/RLlibActionMask-MockUp.html).

From the notebook, you can clearly see that the `tune.run` training goes well as all 3 metrics I care about are rising: `episode_reward_max`, `episode_reward_mean`, `episode_reward_min`. However, when I want to make an inference ( **cell 45** ) with `compute_single_action` or `compute_action`, I am getting very low rewards, even though I do the [checkpointing](https://github.com/ray-project/ray/issues/10290#issuecomment-694440866)… It looks like random, but I really need the very best possible solution to be returned.

Also note that I am new to the RLlib and Ray in general so that I might be missing something “obvious” 🙂

P.S. I am aware of this [issue](https://github.com/ray-project/ray/issues/21417) and related topics here on discuss, but playing with `unsquash_action` and `clip_action` did not help 😕. This was already told to me on [the slack channel](https://ray-distributed.slack.com/archives/CMVUQ22JD/p1652267303922859).

**How severe does this issue affect your experience of using Ray?**

- High: It blocks me from completing my task.

---

<div class="post-metadata">

**Author:** ![arturn](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/arturn/32/2096_2.png) [@arturn](https://discuss.ray.io/u/arturn)\
**Post date:** [May 27, 2022, 3:06pm UTC](https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238/2 "2022-05-27T15:06:36Z")

</div>

Hey @vlainic ,

I can’t say for sure, but here are a couple of things that I find worth looking at:

- Your write that you checkpoint only at the end of training in cell 43, but the code before looks like you checkpoint at the most promising Trainer.step(), call?

- Even though these graphs look cool, it is generally worth mentioning that when checkpointing, you want to have an estimate of your policy’s performance that is as free from variance as possible. So evaluating over multiple episodes would be a better way to choose where to checkpoint.

- One final thought that might not apply to your case: Have a look at your KL loss! Is it down at your checkpoint? If it’s staying up, you have optimized for a stochastic policy and should not disable exploration during evaluation, since this will get rid of stochasticity. That’s a rare case, I just wanted to mention it since I don’t fully understand your environment.

Let us know if this helped!

---

<div class="post-metadata">

**Author:** ![vlainic](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/vlainic/32/2330_2.png) [@vlainic](https://discuss.ray.io/u/vlainic)\
**Post date:** [May 30, 2022, 5:57pm UTC](https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238/3 "2022-05-30T17:57:48Z")

</div>

Hello @arturn,

Thanks for the response. So here is what I understand from you:

- I should have `train_batch_size` much larger than the episode length? For the current mock-up set up `episode_length` is static and it is 28. What about `rollout_fragment_length`?
- How do I evaluate over multiple episodes? Is it with the `checkpoint_freq=n`? I understood that parameter as skipping several `n-1` and then checking the `n-th`, not averaging over the last `n`.
- where can I find the KL loss? Is it `total_loss` or `kl_coeff` or something third?

Today I did play a bit with `train_batch_size=2800` (with significant increase in `episodes_total` too), `rollout_fragment_length: [280,560]`, increased `sgd_minibatch_size`, tried both options for `checkpoint_at_end`, changing `checkpoint_freq` and `keep_checkpoints_num`… but inference is still the same…

Why is a short episode a problem, btw?

---

<div class="post-metadata">

**Author:** ![vlainic](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/vlainic/32/2330_2.png) [@vlainic](https://discuss.ray.io/u/vlainic)\
**Post date:** [June 8, 2022, 8:27am UTC](https://discuss.ray.io/t/bad-inference-after-perfect-training-what-am-i-missing/6238/4 "2022-06-08T08:27:30Z")

</div>

[SOLUTION] I had to change `observations`, i.e. model input.

What I had initially as input was a row from the mask as it changes with action. Example:

- Start (cell 10): `[-1., -1., -1., -1., -1., -1., -1., -1., -1., -1., -1.]`
- cell 48 from now on:
  - Action = 0 → `[1., 0., 0., -1., -1., -1., -1., -1., -1., -1., -1.]`
  - Action = 5 → `[1., 0., 0., 0., 0., 1., -1., -1., -1., -1., -1.]`
  - Action = 6 → `[1., 0., 0., 0., 0., 1., 1., 0., 0., -1., -1.]`
  - Action = 9 → well… you can guess it 🙂

However, the above did not work and what finally works the best is just putting the sequence of the last actions done. So to compare it to the previous example:

- Start: `[-1,-1]`
- Action = 0 → `[0, -1]`
- Action = 5 → `[5, 0]`
- Action = 6: `[6, 5]`
- Action = 9: guess again!

Not sure why it was easier for the model to learn from the latter inputs, but that worked 🎉

P.S. From this, it looks that this would be better to approach with LSTM/RNN models, but the problem above is super-simplified and in reality, besides this action sequence there will be more side inputs. Anyway, this is not the topic of this issue…
