-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsupervised_actor_critic.py
More file actions
200 lines (140 loc) · 7.08 KB
/
Copy pathsupervised_actor_critic.py
File metadata and controls
200 lines (140 loc) · 7.08 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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
"""
訓練データの着手と勝敗から, Actor-Criticの学習則で方策と価値を教師あり学習するためのスクリプト
"""
import multiprocessing as mp
import torch
from torch import nn
from torch.utils.data import DataLoader
from gomoku import Position
from dataset import PositionDataset
from dual_net import DualNet
OUTCOME_WIN = 1
OUTCOME_LOSS = 0
OUTCOME_DRAW = 0.5
class SupervisedActorCriticConfig:
def __init__(self):
self.initial_model_path = None
self.train_dataset_path = "data/train_data.txt"
self.test_dataset_path = "data/test_data.txt"
self.model_out_path = "dualnet.pth"
self.train_loss_history_path = "train_loss_history.txt"
self.test_loss_history_path = "test_loss_history.txt"
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 訓練データと読み込む際のワーカー数.
self.num_workers = 8
self.board_size = 9
self.learning_rate = 0.001
self.l2_penalty = 1.0e-4
self.batch_size = 4096
self.max_epoch = 1000
# テスト損失がpatience回連続して改善しない場合は学習を打ち切る.
self.patience = 10
assert(self.initial_model_path is not None or not self.transfer_learning)
def loss_func(model: DualNet, batch: tuple[torch.Tensor, torch.Tensor, torch.Tensor]) -> tuple[float, float, float]:
"""
モデルの損失を計算する関数.
"""
pos_tensor, move_tensor, outcome_tensor = batch
p_logit, v_logit = model(pos_tensor)
v = nn.functional.sigmoid(v_logit).detach()
value_loss = nn.functional.binary_cross_entropy_with_logits(v_logit.squeeze(1), outcome_tensor)
log_p = nn.functional.log_softmax(p_logit, dim=1)
cross_entropy = nn.functional.nll_loss(log_p, move_tensor, reduction='none')
# ネットワークが予測した価値と実際の勝敗の差(アドバンテージ).
advantage = outcome_tensor - v.squeeze(1)
# クロスエントロピーをアドバンテージで重みづけをする(Actor-Critic).
# これにより負けた対局における着手の確率が低くなる.
# ただし,アドバンテージには引き分けの価値を加算する.
# これは,不利な局面から負けた場合は,必ずしも悪手を打ったとは限らないので,
# その場合の重みを負にしないため.
policy_loss = (cross_entropy * (advantage + OUTCOME_DRAW)).mean()
return policy_loss, value_loss
def model_step(model: DualNet, optimizer, batch: tuple[torch.Tensor, torch.Tensor, torch.Tensor]) -> tuple[float, float]:
"""
モデルの1ステップ分の更新を実行し, 方策と価値の損失を返す.
"""
policy_loss, value_loss = loss_func(model, batch)
loss = value_loss + policy_loss
optimizer.zero_grad()
loss.backward()
optimizer.step()
return policy_loss.item(), value_loss.item()
def evaluate_model(config: SupervisedActorCriticConfig, model: DualNet, dataloader: DataLoader) -> tuple[float, float]:
"""
モデルの評価を行い, 平均方策損失と平均価値損失を返す.
"""
model.eval()
total_policy_loss = 0.0
total_value_loss = 0.0
num_batches = 0
with torch.no_grad():
for batch in dataloader:
batch = tuple(tensor.to(config.device) for tensor in batch)
policy_loss, value_loss = loss_func(model, batch)
total_policy_loss += policy_loss.item()
total_value_loss += value_loss.item()
num_batches += 1
model.train()
return total_policy_loss / num_batches, total_value_loss / num_batches
if __name__ == "__main__":
mp.set_start_method("spawn", force=True)
config = SupervisedActorCriticConfig()
# モデルとオプティマイザの初期化
model = DualNet(config.board_size)
if config.initial_model_path is not None:
model.load_state_dict(torch.load(config.initial_model_path))
if config.transfer_learning:
model.fix_shared_weights()
model.init_action_head_weights()
model.init_value_head_weights()
elif config.initial_model_path is None:
model.unfix_shared_weights()
model.init_weights()
model.to(config.device)
optimizer = torch.optim.Adam(model.parameters(), lr=config.learning_rate, weight_decay=config.l2_penalty)
print("loading dataset...")
train_dataset = PositionDataset(config.train_dataset_path)
test_dataset = PositionDataset(config.test_dataset_path)
train_dataloader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=config.num_workers, pin_memory=True, persistent_workers=True)
test_dataloader = DataLoader(test_dataset, batch_size=config.batch_size, shuffle=False, num_workers=config.num_workers, pin_memory=True, persistent_workers=True)
print(f"the number of training samples: {len(train_dataset)} positions")
print(f"the number of test samples: {len(test_dataset)} positions")
print(f"batch size: {config.batch_size}\n")
print("evaluating initial model...")
policy_loss, value_loss = evaluate_model(config, model, test_dataloader)
print(f"Initial test loss: policy_loss={policy_loss:.4f}, value_loss={value_loss:.4f}\n")
best_test_loss = policy_loss + value_loss
train_loss_history = []
test_loss_history = []
patience_counter = 0
test_loss_history.append((policy_loss, value_loss))
print("start training...")
for epoch in range(config.max_epoch):
print(f"Epoch: [{epoch + 1}/{config.max_epoch}]")
for batch_id, batch in enumerate(train_dataloader):
batch = tuple(tensor.to(config.device) for tensor in batch)
policy_loss, value_loss = model_step(model, optimizer, batch)
train_loss_history.append((policy_loss, value_loss))
if (batch_id + 1) % 100 == 0:
print(f"Batch [{batch_id + 1}/{len(train_dataloader)}]: "
f"policy_loss={policy_loss:.4f}, value_loss={value_loss:.4f}")
# エポックごとにテストデータで評価
policy_loss, value_loss = evaluate_model(config, model, test_dataloader)
test_loss_history.append((policy_loss, value_loss))
print(f"Test loss: policy_loss={policy_loss:.4f}, value_loss={value_loss:.4f}")
current_test_loss = policy_loss + value_loss
if current_test_loss < best_test_loss:
best_test_loss = current_test_loss
patience_counter = 0
# モデルの保存
torch.save(model.state_dict(), config.model_out_path)
print(f"Model saved at epoch {epoch + 1} with test loss {best_test_loss:.4f}")
else:
patience_counter += 1
if patience_counter >= config.patience:
print("Early stopping triggered.")
break
with open(config.train_loss_history_path, 'w') as f:
f.write(str(train_loss_history))
with open(config.test_loss_history_path, 'w') as f:
f.write(str(test_loss_history))