# \[rllib\] Dict Action Space and Custom Model

**URL:** <https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445>\
**Category:** RLlib\
**Created:** [March 29, 2021, 8:31am UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445 "2021-03-29T08:31:24Z")\
**Posts on this page:** 8\
**Page:** 1

<div class="post-metadata">

**Author:** ![Sertingolix](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/sertingolix/32/649_2.png) [@Sertingolix](https://discuss.ray.io/u/Sertingolix)\
**Post date:** [March 29, 2021, 8:31am UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/1 "2021-03-29T08:31:24Z")

</div>

Hi there,

I’m sure it is a rather simple thing, but I did not find an example for it. The default Model can handle Dict action spaces. What I want to know is, how do I output dict action spaces in my custom model?

I tried it with directly outputting the dict, which results in

```auto
File "/home/ray/anaconda3/lib/python3.7/site-packages/ray/rllib/models/torch/torch_action_dist.py", line 386, in __init__
    inputs = torch.from_numpy(inputs)
TypeError: expected np.ndarray (got dict)

```

Which is not surprising, as I in fact outputted a dict. Should I flatten it? Is there already a function in rllib doing this (such that my outputs align) or should one always write a custom action distribution?

Thank you for answering. (Especially if it is obvious to you)

---

<div class="post-metadata">

**Author:** ![RickLan](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/ricklan/32/901_2.png) [@RickLan](https://discuss.ray.io/u/RickLan)\
**Post date:** [March 29, 2021, 9:35am UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/2 "2021-03-29T09:35:00Z")

</div>

Would you mind sharing the code of your custom model?

---

<div class="post-metadata">

**Author:** ![Sertingolix](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/sertingolix/32/649_2.png) [@Sertingolix](https://discuss.ray.io/u/Sertingolix)\
**Post date:** [March 29, 2021, 10:55am UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/3 "2021-03-29T10:55:36Z")

</div>

sure. Maybe you can spot a mistake in my thought process.

I know that the expected number of outputs should be 64. I do not get how to transform my dict into these dimensions. I assume it should be flattened as the flattened dimension of my dict would be 64 but i want to make sure, this is correct

```auto
import torch
from torch import nn
import numpy as np

from .._net_utils.task import BaseTask

import ray
from ray.rllib.models.torch.torch_modelv2 import TorchModelV2
from ray.rllib.models.torch.fcnet import FullyConnectedNetwork

class FloodingTorchModel(TorchModelV2, nn.Module):
    def __init__ (self, obs_space, action_space, num_outputs, model_config, name):
        print("-------building model")
        print("obs_space", obs_space.original_space)
        print("action_space", action_space)
        print("num_outputs", num_outputs)
        print("model_config", model_config)
        super(). __init__ (obs_space, action_space, num_outputs, model_config,name)
        nn.Module. __init__ (self)

        # get properties/dimension of current task
        self.task = model_config.get("task", BaseTask())

        self.flat_space = obs_space
        self.obs_space= obs_space.original_space
        self.action_space = action_space

        self.value = None

        self.relu = nn.ReLU()
        self.softmax = nn.Softmax(dim=2)

        self.task.sens_space.space

        self.hidden_dim=256

        self.tohidden = nn.Linear(self.flat_space.shape[0],self.hidden_dim)
        self.hidden_layer = nn.Linear(self.hidden_dim,self.hidden_dim)

        self.hidden2node_state = nn.Linear(self.hidden_dim,self.task.node_emb_space.length)
        self.hidden2action = nn.Linear(self.hidden_dim,2)
        self.hidden2send = nn.Linear(self.hidden_dim,2)
        self.hidden2msg = nn.Linear(self.hidden_dim,4)
        self.hidden2value = nn.Linear(self.hidden_dim,1)

    def forward(self, input_dict, state, seq_lens):
        obs = input_dict["obs_flat"]
        # Store last batch size for value_function output.
        self._last_batch_size = obs.shape[0]
        
        # to use original observation space

        # observation = input_dict["obs"]

        # sensor = observation["sensor"]
        # node_state = observation["node_state"]
        # edge_emb = observation["edge_emb"]
        # msg = observation["msg"]
        # ids = observation["id"]
        # time = observation["time"]

        # print("sensor shape", sensor.shape)
        # print("node_state shape", node_state.shape)
        # print("msg shape", len(msg),msg[0].shape)

        hidden = self.tohidden(obs)
        hidden = self.relu(hidden)
        hidden = self.hidden_layer(hidden)

        node_state = self.hidden2node_state(hidden)
        action = self.hidden2action(hidden)

        send_msg = self.hidden2send(hidden)
        msg = self.hidden2msg(hidden)

        send_msg = torch.unsqueeze(send_msg,1).expand(-1,self.task.max_degree,-1)
        msg = torch.unsqueeze(msg,1).expand(-1,self.task.max_degree,-1)

        action = {
            'node_state': node_state,
            'action': action,
            'send_msg': send_msg,
            'msg': msg,
        }

        print("-----action shape")
        for key,val in action.items():
            print(f"key: {key} value_shape: {val.shape}")

        self.value = self.hidden2value(hidden)

        #the action should be processed/ changed
        #how do i know the order of the dict when flattening?
        return action, state

    def value_function(self):
        return self.value

```

Thank you very much

---

<div class="post-metadata">

**Author:** ![RickLan](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/ricklan/32/901_2.png) [@RickLan](https://discuss.ray.io/u/RickLan)\
**Post date:** [March 29, 2021, 4:23pm UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/4 "2021-03-29T16:23:09Z")

</div>

> [@Sertingolix](#):
>
> I know that the expected number of outputs should be 64

How do you know this? Is action\_space defined on your environment?

I don’t know enough about RLlib to know if dict-based action space is supported. However my intuition is that the output of forward() needs to be torch variables because they are used in the loss function calculation and the gradient tape. So if you are using dict, then you probably need a custom loss function.

---

<div class="post-metadata">

**Author:** ![Sertingolix](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/sertingolix/32/649_2.png) [@Sertingolix](https://discuss.ray.io/u/Sertingolix)\
**Post date:** [March 30, 2021, 9:30am UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/5 "2021-03-30T09:30:43Z")

</div>

> [@RickLan](#):
>
> How do you know this? Is action\_space defined on your environment?

Yes, the action space is defined in my environment as follows (not that it really matters):

#logits a 2, c 2, (each tuple space is 10 long) b 10 \* 4, d 10 \* 2

```auto
Dict(a:Discrete(2), b:Tuple(Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4), Discrete(4)), c:Discrete(2), d:Tuple(Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2), Discrete(2)))

```

To fix it i now use the following helper method

```auto
class FloodingTorchModel(TorchModelV2, nn.Module):
    def __init__ (self, obs_space, action_space, num_outputs, model_config, name):
    ....
    self.action_space = action_space
    self.order = list(sorted(self.action_space.spaces.keys()))

def forward(self, input_dict, state, seq_lens):
   ...
        action = self.dict_action2preprocessed_action(action)
   ...
def dict_action2preprocessed_action(self,action):
        """ 
        The action has to be a flatended tensor 
        In order to make sure the flattened elements allign we use this converter
        """

        stack = []
        for key in self.order:
            stack.append(torch.flatten(action[key], start_dim=1))

        cat = torch.cat(stack,1)
        print("cat shape ", cat.shape)
        return cat

```

One should note, that rllib at the moment alphabetically orders dict keys.

> [@RickLan](#):
>
> I don’t know enough about RLlib to know if dict-based action space is supported. However my intuition is that the output of forward() needs to be torch variables because they are used in the loss function calculation and the gradient tape. So if you are using dict, then you probably need a custom loss function.

You are right, [BATCH, num\_outputs] is expected from forward. It seems to run without a custom loss function.

Thank you for your help

---

<div class="post-metadata">

**Author:** ![sven1977](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/sven1977/32/53_2.png) [@sven1977](https://discuss.ray.io/u/sven1977)\
**Post date:** [March 30, 2021, 3:03pm UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/6 "2021-03-30T15:03:59Z")

</div>

When using dict action spaces, your model should output a flat tensor, which will then be passed into a MultiActionDistribution for action sampling. This sampling step then returns a dict.  
The alphabetic sorting is potentially a problem, however, it’s forced upon RLlib via gym’s very own Dict space handling (`Dict.spaces` is an `OrderedDict`).

If you check the code in MultiActionDistribution ([ray/torch\_action\_dist.py at master · ray-project/ray · GitHub](https://github.com/ray-project/ray/blob/master/rllib/models/torch/torch_action_dist.py#L391)), you will see that we create an alphabetically sorted `action_space_struct` dict, which we then use to regenerate the action dict from your flat tensor outputs.

In other words, as long as you return from your model a tensor that is sorted alphabetically according your dict (print out `self.action_space_struct` in the MultiActionDistribution to see what the exact order should be in case you have additional nesting going on), it’ll be fine.  
Alternatively, you can use a custom action distribution, which then would handle your model’s output (whatever that would be, e.g. a dict), but then you would be responsible for the “handover” between model and action distribution.

---

<div class="post-metadata">

**Author:** ![buja26](https://avatars.discourse-cdn.com/v4/letter/b/977dab/32.png) [@buja26](https://discuss.ray.io/u/buja26)\
**Post date:** [November 17, 2025, 12:32pm UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/7 "2025-11-17T12:32:13Z")

</div>

Hi there,

I’m having troubles using TorchMultiDistribution since MultiActionDistribution is the old API, I guess. According to you, my RLModule returns a flatten tensor that is sorted alphabetically according to my dict action space. But i get the following error:

> TypeError: TorchMultiDistribution.from\_logits() missing 2 required positional arguments: ‘child\_distribution\_cls\_struct’ and ‘input\_lens’

Here is my code snippet from RLModule:  
`def _pi_outputs(self, z: torch.Tensor) -> Dict[str, torch.Tensor]:`  
` out: Dict[str, torch.Tensor] = {}`

` # Policy Head`  
` outs = []`  
` for k, subspace in self.action_subspaces.items():`  
` head = self.heads[k]`  
` if isinstance(subspace, Discrete):`  
` logits = head(z)`  
` outs.append(logits)`  
` elif isinstance(subspace, Box):`  
` mu = head(z)`  
` log_std = torch.clamp(self.log_std_params[k], self.pol_cfg.log_std_min, self.pol_cfg.log_std_max)`  
` outs.append(torch.cat([mu, log_std.expand_as(mu)], dim=-1))`

` action_logits = torch.cat(outs, dim=-1)`  
` out[Columns.ACTION_DIST_INPUTS] = action_logits`  
` return out`

And I simply train the PPO with:

`algo_config = agent.algo_config`  
`algo = algo_config.build_algo()`  
`algo.train()`

I set `action_dist_cls` to `TorchMultiDistribution`. I am not really sure how to use this distribution since `from_logits` always takes the logits only and not the other additional arguments.

Thank you very much for your help!

---

<div class="post-metadata">

**Author:** ![DenBuzz](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/denbuzz/32/6496_2.png) [@DenBuzz](https://discuss.ray.io/u/DenBuzz)\
**Post date:** [December 1, 2025, 9:56pm UTC](https://discuss.ray.io/t/rllib-dict-action-space-and-custom-model/1445/8 "2025-12-01T21:56:21Z")

</div>

You need to create a distribution class that knows what the child distribution classes are. The parent class of the TorchMultiDistribution (Distribution) has a helper method for this called `get_partial_dist_cls`. You can see the source here: [ray/rllib/core/distribution/distribution.py at ff0cd7f257bcdc6b6129d622dc4536bb20d3188a · ray-project/ray · GitHub](https://github.com/ray-project/ray/blob/ff0cd7f257bcdc6b6129d622dc4536bb20d3188a/rllib/core/distribution/distribution.py#L189)

I believe you just need to use it like this:

`self.action_dist_cls = TorchMultiDistribution.get_partial_dist_cls(child_distribution_struct={"foo": Discrete(5), "bar": Discrete(3)})`

Or something along those lines. The resulting distribution class then knows what the underlying struct is.

Hope that helps!
