AI4Finance-Foundation / AI4Finance-Foundation/ElegantRL

🐛 ✨ the update methods of std for running stat is not good

オープン
#258 コメント 0 件 リアクション 0 件 担当者 3 名 @shixun404 が担当を希望しています GitHub で見る
refactoring
主要言語
Python
スター
4.4k
フォーク
978
PR マージ指標
30日以内にマージされた PR はありません

説明

The better way to get the running stat of std:
```
def update_avg_std_for_state_value_norm(self, states: Tensor, returns: Tensor):
tau = self.state_value_tau

state_avg = states.mean(dim=0, keepdim=True)
state_vam = (states ** 2).mean(dim=0, keepdim=True)
self.cri.state_avg[:] = self.cri.state_avg * (1 - tau) + state_avg * tau
self.cri.state_vam[:] = self.cri.state_vam * (1 - tau) + state_vam * tau
self.cri.state_std[:] = torch.sqrt(self.cri.state_vam - self.cri.state_avg ** 2) + 1e-4

returns_avg = returns.mean(dim=0)
returns_vam = (returns ** 2).mean(dim=0)
self.cri.value_avg[:] = self.cri.value_avg * (1 - tau) + returns_avg * tau
self.cri.value_vam[:] = self.cri.value_vam * (1 - tau) + returns_vam * tau
self.cri.value_std[:] = torch.sqrt(self.cri.value_vam - self.cri.value_avg ** 2) + 1e-4

self.act.state_avg[:] = self.cri.state_avg
self.act.state_std[:] = self.cri.state_std
```

It is better than the current way of ElegantRL:

https://github.com/AI4Finance-Foundation/ElegantRL/blob/49f9d8ca403e146f00181c74f849f847c5d7aae7/elegantrl/agents/base.py#L241-L256

---

By the way, we should change the `std` to `vam` in `QNetBase`, `ActorBase` and `CriticBase`.
```
class CriticBase(nn.Module):
def __init__(self, state_dim: int, action_dim: int):
super().__init__()
self.state_dim = state_dim
self.action_dim = action_dim
self.net = None # build_mlp(dims=[state_dim + action_dim, *dims, 1])

self.state_avg = nn.Parameter(torch.zeros((state_dim,)), requires_grad=False)
self.state_vam = nn.Parameter(torch.ones((state_dim,)), requires_grad=False) # var.mean
self.state_std = nn.Parameter(torch.ones((state_dim,)), requires_grad=False) # sqrt(var.mean - avg.pow2)

self.value_avg = nn.Parameter(torch.zeros((1,)), requires_grad=False)
self.value_vam = nn.Parameter(torch.ones((1,)), requires_grad=False) # var.mean
self.value_std = nn.Parameter(torch.ones((1,)), requires_grad=False) # sqrt(var.mean - avg.pow2)

def state_norm(self, state: Tensor) -> Tensor:
return (state - self.state_avg) / self.state_std

def value_re_norm(self, value: Tensor) -> Tensor:
return value * self.value_std + self.value_avg
```

https://github.com/AI4Finance-Foundation/ElegantRL/blob/49f9d8ca403e146f00181c74f849f847c5d7aae7/elegantrl/agents/net.py#L321-L336

コントリビューションガイド

このリポジトリのコントリビューションガイドは索引されていません

評価

この issue はまだ評価されていません。

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。