AI4Finance-Foundation / AI4Finance-Foundation/ElegantRL

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

Aperta
#258 0 commenti 0 reazioni 3 assegnatari Rivendicata da @shixun404 Vedi su GitHub
refactoring
Lingua principale
Python
Stelle
4.4k
Fork
978
Metriche di merge delle PR
Nessuna PR unita negli ultimi 30g

Descrizione

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

Guida per i contributori

Nessuna guida per i contributori indicizzata per questo repository

Valutazione

Questa issue non è ancora stata valutata.

Ricevi le nuove issue nella tua casella

Un breve riepilogo di issue GitHub adatte ai principianti.