# \[RLlib\] Exporting a PyTorch policy for TorchScript

**URL:** <https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314>\
**Category:** RLlib\
**Created:** [December 23, 2020, 12:16am UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314 "2020-12-23T00:16:39Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![amcleod](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/amcleod/32/190_2.png) [@amcleod](https://discuss.ray.io/u/amcleod)\
**Post date:** [December 23, 2020, 12:16am UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314/1 "2020-12-23T00:16:39Z")

</div>

Hello all,

I’m trying to use Ray/RLLib to train a policy, and I’m running into trouble when it comes time to export it. I’m using Ray 1.0.1.post1. My workflow uses the Python API and a custom environment, and iteratively calling the train() method on a Trainer object. After it completes, I have a model buried somewhere in the trainer object. Here’s what I’ve done to try to free it:

- If I call `trainer.export_model(ExportFormat.MODEL, args.export_dir)` I get a NotImplementedError:

> File “/home/amcleod/.local/lib/python3.6/site-packages/ray/rllib/policy/torch\_policy.py”, line 590, in export\_model  
> raise NotImplementedError

- I can get a policy object with `trainer.get_policy()`, but this object can’t be pickled.

Ideally, I would like to find a way to end up with a torch module containing the trained policy exclusively, that is free of the Ray class hierarchy. How would I go about doing this?

Thanks.

---

<div class="post-metadata">

**Author:** ![Eric\_Adlam](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/eric_adlam/32/394_2.png) [@Eric\_Adlam](https://discuss.ray.io/u/Eric_Adlam)\
**Post date:** [February 5, 2021, 7:40pm UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314/2 "2021-02-05T19:40:19Z")

</div>

Did you ever figure out how to do this?

---

<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:** [February 8, 2021, 4:35pm UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314/3 "2021-02-08T16:35:38Z")

</div>

To implement this missing method for TorchPolicy is harder than I hoped.

The seemingly simple solution is to add this to the TorchPolicy class:

```auto
    @override(Policy)
    @DeveloperAPI
    def export_model(self, export_dir: str) -> None:
        """Exports the Policy's Model to local directory for serving.

        Creates a TorchScript model and saves it.

        Args:
            export_dir (str): Local writable directory or filename.
        """
        dummy_inputs = self._lazy_tensor_dict(self._dummy_batch.data)
        # Provide dummy state inputs if not an RNN (torch cannot jit with empty list).
        if "state_in_0" not in dummy_inputs:
            dummy_inputs["state_in_0"] = dummy_inputs["seq_lens"] = np.array([1.0])
        dummy_inputs = {k: dummy_inputs[k] for k in dummy_inputs.keys()}
        traced = torch.jit.trace(self.model, dummy_inputs)
        if os.path.isfile(export_dir):
            file_name = export_dir
        else:
            file_name = os.path.join("model.pt", export_dir)
        traced.save(file_name)

```

However, torch jit requires the all return values of the nn.Module to be tensors, which is not the case for our TorchModelV2, which - if not an RNN - returns an empty list of internal states (as second return value). Removing this for non-RNNs would break our entire ModelV2 API. 😕

---

<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:** [February 8, 2021, 4:48pm UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314/4 "2021-02-08T16:48:32Z")

</div>

Hmm, I actually did find a way, but it would require you to pass in a fake state\_in\_0 (not `[]`!) and seq\_lens tensor (not None!). I guess this is better than not having this work at all.  
We may have to change the Model API (or provide an alternative one) at some point to make this work properly.  
I’ll PR. …

---

<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:** [February 8, 2021, 6:31pm UTC](https://discuss.ray.io/t/rllib-exporting-a-pytorch-policy-for-torchscript/314/5 "2021-02-08T18:31:28Z")

</div>

> <https://github.com/ray-project/ray/pull/13989>
