-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_mc.py
More file actions
87 lines (63 loc) · 2.51 KB
/
Copy pathtrain_mc.py
File metadata and controls
87 lines (63 loc) · 2.51 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""
뚱이 게임 Monte Carlo(에피소드 기반) 정책 gradient 학습 스크립트.
- 알고리즘: REINFORCE (Monte Carlo Policy Gradient)
- 환경: DdungEnv (Maskable 액션 사용)
- 목적: PPO / DQN 이외에, 에피소드 전체 리턴을 사용하는 MC 계열 알고리즘 실험용
의존성 설치: pip install -r requirements.txt
실행: python train_mc.py
학습된 정책 파라미터는 models/ddung_mc_policy.pt 로 저장됩니다.
"""
import os
from typing import List
from tqdm import tqdm
import torch
import torch.optim as optim
from ddung_env import DdungEnv
from mc_policy import PolicyNet, select_action
def main() -> None:
os.makedirs("models", exist_ok=True)
save_path = os.path.join("models", "ddung_mc_policy.pt")
env = DdungEnv(time_limit_sec=120)
obs_dim = env.observation_space.shape[0]
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
policy = PolicyNet(obs_dim).to(device)
optimizer = optim.Adam(policy.parameters(), lr=1e-3)
gamma = 0.99
num_episodes = 3_000 # 실험/공부용 에피소드 수
pbar = tqdm(range(1, num_episodes + 1), desc="Train MC", unit="ep")
for ep in pbar:
obs, _ = env.reset()
log_probs: List[torch.Tensor] = []
rewards: List[float] = []
while True:
action_mask = env.action_masks()
action, log_prob = select_action(policy, obs, action_mask, device)
obs, reward, terminated, truncated, info = env.step(action)
log_probs.append(log_prob)
rewards.append(reward)
if terminated or truncated:
break
# Monte Carlo 리턴 계산 (에피소드 전체)
returns: List[float] = []
G = 0.0
for r in reversed(rewards):
G = r + gamma * G
returns.append(G)
returns.reverse()
returns_t = torch.as_tensor(returns, dtype=torch.float32, device=device)
# 간단한 정규화 (안정성용)
if len(returns_t) > 1:
returns_t = (returns_t - returns_t.mean()) / (returns_t.std() + 1e-8)
log_probs_t = torch.stack(log_probs)
loss = -(log_probs_t * returns_t).sum()
optimizer.zero_grad()
loss.backward()
optimizer.step()
if ep % 100 == 0:
total_return = sum(rewards)
pbar.set_postfix(ret=f"{total_return:.1f}")
torch.save(policy.state_dict(), save_path)
env.close()
print(f"MC policy saved to {save_path}")
if __name__ == "__main__":
main()