jongchullee commited on
Commit
b0759ac
ยท
verified ยท
1 Parent(s): f87fa76

Upload dqn_agent.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. dqn_agent.py +196 -0
dqn_agent.py ADDED
@@ -0,0 +1,196 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import random
2
+ from collections import deque
3
+ import numpy as np
4
+ import torch
5
+ import torch.nn as nn
6
+ import torch.optim as optim
7
+
8
+ class DuelingDQN(nn.Module):
9
+ """
10
+ Dueling DQN ๊ตฌ์กฐ:
11
+ ๊ฐ€์น˜(Value) ํ•จ์ˆ˜์™€ ์ด์ (Advantage) ํ•จ์ˆ˜๋ฅผ ๋ถ„๋ฆฌํ•˜์—ฌ ๋” ๋น ๋ฅด๊ณ  ์•ˆ์ •์ ์ธ Q-ํ•™์Šต ์ œ๊ณต
12
+ Q(s, a) = V(s) + (A(s, a) - mean(A(s)))
13
+ """
14
+ def __init__(self, state_dim=8, action_dim=4):
15
+ super(DuelingDQN, self).__init__()
16
+
17
+ # ๊ณตํ†ต ํŠน์ง• ์ถ”์ถœ์ธต
18
+ self.feature_network = nn.Sequential(
19
+ nn.Linear(state_dim, 128),
20
+ nn.ReLU(),
21
+ nn.Linear(128, 128),
22
+ nn.ReLU()
23
+ )
24
+
25
+ # ์ƒํƒœ ๊ฐ€์น˜(State Value) ์ŠคํŠธ๋ฆผ V(s)
26
+ self.value_stream = nn.Sequential(
27
+ nn.Linear(128, 64),
28
+ nn.ReLU(),
29
+ nn.Linear(64, 1)
30
+ )
31
+
32
+ # ํ–‰๋™ ์ด์ (Action Advantage) ์ŠคํŠธ๋ฆผ A(s, a)
33
+ self.advantage_stream = nn.Sequential(
34
+ nn.Linear(128, 64),
35
+ nn.ReLU(),
36
+ nn.Linear(64, action_dim)
37
+ )
38
+
39
+ def forward(self, state):
40
+ features = self.feature_network(state)
41
+ values = self.value_stream(features)
42
+ advantages = self.advantage_stream(features)
43
+
44
+ # Dueling ๊ณต์‹: Q = V + (A - mean(A))
45
+ q_values = values + (advantages - advantages.mean(dim=-1, keepdim=True))
46
+ return q_values
47
+
48
+
49
+ class ReplayBuffer:
50
+ """๊ฒฝํ—˜ ๋ฆฌํ”Œ๋ ˆ์ด ๋ฒ„ํผ"""
51
+ def __init__(self, capacity=100000):
52
+ self.buffer = deque(maxlen=capacity)
53
+
54
+ def push(self, state, action, reward, next_state, done):
55
+ self.buffer.append((state, action, reward, next_state, done))
56
+
57
+ def sample(self, batch_size):
58
+ batch = random.sample(self.buffer, batch_size)
59
+ states, actions, rewards, next_states, dones = zip(*batch)
60
+ return (
61
+ np.array(states, dtype=np.float32),
62
+ np.array(actions, dtype=np.int64),
63
+ np.array(rewards, dtype=np.float32),
64
+ np.array(next_states, dtype=np.float32),
65
+ np.array(dones, dtype=np.float32),
66
+ )
67
+
68
+ def __len__(self):
69
+ return len(self.buffer)
70
+
71
+
72
+ class DQNAgent:
73
+ """
74
+ LunarLander ์ฐฉ๋ฅ™์šฉ Dueling Double-DQN Agent
75
+ - Double DQN: ์˜ค๋ฒ„์—์Šคํ‹ฐ๋ฉ”์ด์…˜ ๋ฐฉ์ง€
76
+ - Epsilon Decay: 100% -> 5% ์ ์ง„์  ๊ฐ์‡„
77
+ - Soft Target Update (Polyak averaging)
78
+ """
79
+ def __init__(
80
+ self,
81
+ state_dim=8,
82
+ action_dim=4,
83
+ lr=5e-4,
84
+ gamma=0.99,
85
+ tau=0.005,
86
+ buffer_size=100000,
87
+ batch_size=64,
88
+ epsilon_start=1.0,
89
+ epsilon_end=0.05,
90
+ total_episodes=1000
91
+ ):
92
+ self.state_dim = state_dim
93
+ self.action_dim = action_dim
94
+ self.gamma = gamma
95
+ self.tau = tau
96
+ self.batch_size = batch_size
97
+
98
+ # ํƒ์ƒ‰๋ฅ (Epsilon) ํŒŒ๋ผ๋ฏธํ„ฐ (1.0 -> 0.05)
99
+ self.epsilon_start = epsilon_start
100
+ self.epsilon_end = epsilon_end
101
+ self.total_episodes = total_episodes
102
+ self.epsilon = epsilon_start
103
+
104
+ # 1000 ์—ํ”ผ์†Œ๋“œ ์ค‘ ์•ฝ 75% ์ง€์ (750 ep)์—์„œ epsilon_end(0.05)์— ๋„๋‹ฌํ•˜๋„๋ก ๊ฐ์‡„์œจ ์„ค์ •
105
+ self.epsilon_decay = (epsilon_end / epsilon_start) ** (1.0 / (total_episodes * 0.75))
106
+
107
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
108
+
109
+ # ์ •์ฑ…๋ง & ํƒ€๊นƒ๋ง ์ƒ์„ฑ
110
+ self.policy_net = DuelingDQN(state_dim, action_dim).to(self.device)
111
+ self.target_net = DuelingDQN(state_dim, action_dim).to(self.device)
112
+ self.target_net.load_state_dict(self.policy_net.state_dict())
113
+ self.target_net.eval()
114
+
115
+ self.optimizer = optim.AdamW(self.policy_net.parameters(), lr=lr, weight_decay=1e-4)
116
+ self.criterion = nn.SmoothL1Loss() # Huber Loss (๋…ธ์ด์ฆˆ์— ๊ฐ•๊ฑด)
117
+ self.memory = ReplayBuffer(buffer_size)
118
+
119
+ def select_action(self, state, evaluate=False):
120
+ """
121
+ ํ–‰๋™ ์„ ํƒ ๋ฐ ๊ฐ ํ–‰๋™๋ณ„ Q-Value ๋ฐ˜ํ™˜ (์›น ๋Œ€์‹œ๋ณด๋“œ ์‹œ๊ฐํ™”์šฉ)
122
+ evaluate=True ์ผ ๊ฒฝ์šฐ ์ˆœ์ˆ˜ Greedy ํ–‰๋™ (๊ฐ„์ง€ ์ฐฉ๋ฅ™ ์‹œ์—ฐ)
123
+ """
124
+ state_t = torch.FloatTensor(state).unsqueeze(0).to(self.device)
125
+
126
+ with torch.no_grad():
127
+ q_values = self.policy_net(state_t).cpu().numpy()[0]
128
+
129
+ if not evaluate and random.random() < self.epsilon:
130
+ action = random.randrange(self.action_dim)
131
+ else:
132
+ action = int(np.argmax(q_values))
133
+
134
+ return action, q_values.tolist()
135
+
136
+ def update(self):
137
+ """Double DQN ๊ธฐ๋ฐ˜ ์‹ ๊ฒฝ๋ง ๊ฐ€์ค‘์น˜ 1์Šคํ… ์—…๋ฐ์ดํŠธ"""
138
+ if len(self.memory) < self.batch_size:
139
+ return None
140
+
141
+ states, actions, rewards, next_states, dones = self.memory.sample(self.batch_size)
142
+
143
+ states_t = torch.FloatTensor(states).to(self.device)
144
+ actions_t = torch.LongTensor(actions).unsqueeze(1).to(self.device)
145
+ rewards_t = torch.FloatTensor(rewards).unsqueeze(1).to(self.device)
146
+ next_states_t = torch.FloatTensor(next_states).to(self.device)
147
+ dones_t = torch.FloatTensor(dones).unsqueeze(1).to(self.device)
148
+
149
+ # ํ˜„์žฌ ์ƒํƒœ์˜ Q-๊ฐ’ ๊ณ„์‚ฐ: Q(s, a)
150
+ curr_q = self.policy_net(states_t).gather(1, actions_t)
151
+
152
+ # Double DQN: Policy Net์œผ๋กœ ์ตœ์  ํ–‰๋™ ์„ ํƒ -> Target Net์œผ๋กœ ํ•ด๋‹น ํ–‰๋™์˜ Q-๊ฐ’ ํ‰๊ฐ€
153
+ with torch.no_grad():
154
+ best_actions = self.policy_net(next_states_t).argmax(1, keepdim=True)
155
+ next_q = self.target_net(next_states_t).gather(1, best_actions)
156
+ target_q = rewards_t + (1.0 - dones_t) * self.gamma * next_q
157
+
158
+ # Loss ๊ณ„์‚ฐ ๋ฐ ์—ญ์ „ํŒŒ
159
+ loss = self.criterion(curr_q, target_q)
160
+
161
+ self.optimizer.zero_grad()
162
+ loss.backward()
163
+ # ์•ˆ์ •์ ์ธ ํ•™์Šต์„ ์œ„ํ•œ Gradient Clipping
164
+ nn.utils.clip_grad_norm_(self.policy_net.parameters(), max_norm=10.0)
165
+ self.optimizer.step()
166
+
167
+ # ํƒ€๊นƒ๋ง Soft Update
168
+ self.soft_update()
169
+
170
+ return loss.item()
171
+
172
+ def soft_update(self):
173
+ """Polyak Averaging ํƒ€๊นƒ๋ง ์†Œํ”„ํŠธ ์—…๋ฐ์ดํŠธ: ฮธ_target = ฯ„*ฮธ_local + (1 - ฯ„)*ฮธ_target"""
174
+ for target_param, policy_param in zip(self.target_net.parameters(), self.policy_net.parameters()):
175
+ target_param.data.copy_(self.tau * policy_param.data + (1.0 - self.tau) * target_param.data)
176
+
177
+ def decay_epsilon(self):
178
+ """์—ํ”ผ์†Œ๋“œ ์ข…๋ฃŒ ์‹œ Epsilon ๊ฐ์‡„"""
179
+ self.epsilon = max(self.epsilon_end, self.epsilon * self.epsilon_decay)
180
+
181
+ def save(self, filepath="best_lunar_lander_dqn.pth"):
182
+ torch.save({
183
+ 'policy_net': self.policy_net.state_dict(),
184
+ 'target_net': self.target_net.state_dict(),
185
+ 'optimizer': self.optimizer.state_dict(),
186
+ 'epsilon': self.epsilon
187
+ }, filepath)
188
+
189
+ def load(self, filepath="best_lunar_lander_dqn.pth"):
190
+ checkpoint = torch.load(filepath, map_location=self.device)
191
+ self.policy_net.load_state_dict(checkpoint['policy_net'])
192
+ self.target_net.load_state_dict(checkpoint['target_net'])
193
+ if 'optimizer' in checkpoint:
194
+ self.optimizer.load_state_dict(checkpoint['optimizer'])
195
+ if 'epsilon' in checkpoint:
196
+ self.epsilon = checkpoint['epsilon']