AI4Finance-Foundation / AI4Finance-Foundation/ElegantRL
🐛 ✨ the update methods of std for running stat is not good
- 主要言語
- 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 はまだ評価されていません。