AI4Finance-Foundation / AI4Finance-Foundation/ElegantRL

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

Abierto
#258 0 comentarios 0 reacciones 3 asignados Reclamado por @shixun404 Ver en GitHub
refactoring
Lenguaje dominante
Python
Estrellas
4.4k
Forks
978
Métricas de merge de PR
Sin PR fusionados en 30 d

Descripción

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

Guía de contribución

No hay ninguna guía de contribución indexada para este repositorio

Evaluación

Este issue todavía no se ha evaluado.

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.