# Partial freeze and partial train

**URL:** <https://discuss.ray.io/t/partial-freeze-and-partial-train/4131>\
**Category:** RLlib\
**Created:** [November 14, 2021, 8:06am UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131 "2021-11-14T08:06:57Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![yiwc](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/yiwc/32/1665_2.png) [@yiwc](https://discuss.ray.io/u/yiwc)\
**Post date:** [November 14, 2021, 8:06am UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/1 "2021-11-14T08:06:57Z")

</div>

Dear RLLib team,

Thanks for your great work in RLLib, we all much enjoy it! However, we met a problem in how to realize this feature, we appreciate it if can help to advise us of your suggested solution in RLLib.

Problem Description:  
This problem requires one more layer of transfer learning on top of another trained policy. The trained policy should be frozen during training and the new layer will be trained.

Our Naive Solution:  
Our plan is to use your custom\_train\_workflow related features. During custom training workflow, we select the part of the model’s parameters to freeze, and some other parts of parameters to train.

Thanks for your help let us know if our solution is correct and follows the rllib philosophy. Or you already have a much easier solution for that.

---

<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:** [November 14, 2021, 11:02pm UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/2 "2021-11-14T23:02:20Z")

</div>

Hi @yiwc ,

If you implement your model yourself, i.e. with the [ModelV2 API](https://github.com/ray-project/ray/blob/0f57a9a105d593b57509d2f48346238259a2d942/rllib/models/modelv2.py#L26) , you can simply put a [tf.stop\_gradient()](https://www.tensorflow.org/api_docs/python/tf/stop_gradient) in your forward pass function.

Otherwise, you can [update](https://github.com/ray-project/ray/blob/e6ae08f41674d2ac1423f3c2a4f8d8bd3500379a/rllib/policy/policy_template.py#L389) your policy with a new `apply_gradients_fn` that only applies the gradients to your one layer and leaves the other ones alone.

If you have questions on how to do this, I will be happy to answer them.

Cheers

---

<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:** [November 15, 2021, 1:58pm UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/3 "2021-11-15T13:58:01Z")

</div>

Hi @yiwc,

If you are using torch, you could write a function to freeze the layers of interest by setting requires\_grad=False on those layers parameters.

If you are using tf and keras you can set the layer.trainable=False

With either framework, could then create a new trainer that apply this function similar to the method described in this post:

> [@How to get mode summary if I use tune.run()?](https://discuss.ray.io/t/how-to-get-mode-summary-if-i-use-tune-run/2024/6):
>
> @bug404, You could also do something like this. I don’t know if it counts as “a better way” but it would work. I do not think that updating after\_init on the Trainer will step on any other callbacks because I think those are defined in the policies rather than the trainer but @sven1977 would know better. from ray.tune.registry import register\_trainable from ray.rllib.agents.ppo import PPOTrainer ModelPrintPPOTrainer = PPOTrainer.with\_updates(after\_init=lambda trainer: trainer.get\_p…

---

<div class="post-metadata">

**Author:** ![yiwc](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/yiwc/32/1665_2.png) [@yiwc](https://discuss.ray.io/u/yiwc)\
**Post date:** [November 15, 2021, 2:42pm UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/4 "2021-11-15T14:42:02Z")

</div>

Hi @arturn @mannyv ,

We appreciate your immediate reply!

Yes we will try apply\_gradients\_fn, see how far we can go.

Thanks again, and have a good day!

---

<div class="post-metadata">

**Author:** ![yiwc](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/yiwc/32/1665_2.png) [@yiwc](https://discuss.ray.io/u/yiwc)\
**Post date:** [November 21, 2021, 4:00am UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/5 "2021-11-21T04:00:14Z")

</div>

Hi @arturn,

Thanks for your advice. Now we are trying to fine-tune a model from a loaded trained model. Where do you think we should put the load model code in?

we thought of a few possible solutions

1. put in the custom train execution plan, before train we load the pre-trained model first.

```auto
# some brief pseudocode just for idea
def execution_plan(xxx):
    policy.load(pretrained_model1,pretrained_model2)

```

1. we can also load the pre-trained model before the tune function starts.

```auto
# some brief pseudocode just for idea
my_trainer=trainer(xxx)
my_trainer.policy.load_models(pretrained_model_weights1,pretrained_model_weights2)

```

Appreciate your help and advice if these are recommended solutions~  
Regards,

---

<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:** [November 21, 2021, 2:47pm UTC](https://discuss.ray.io/t/partial-freeze-and-partial-train/4131/6 "2021-11-21T14:47:13Z")

</div>

Hi @yiwc ,

Glad we could help.  
If you want to load a complete model of a previously trained policy, the easiest way is to call the `restore` method of your Trainer. From the [docs](https://docs.ray.io/en/latest/rllib-training.html):

```auto
agent = ppo.PPOTrainer(config=config, env=env_class)
agent.restore(checkpoint_path)

```

Does this work for you? There are of course other ways and more elaborate solutions.  
I am sure @mannyv has more to offer 🙂

Cheers
