ffzeroHua commited on
Commit
5640999
·
verified ·
1 Parent(s): 3d5a0b0

Upload 2 files

Browse files
Files changed (2) hide show
  1. model3pLOCAL.py +452 -0
  2. model3pNEW.py +418 -0
model3pLOCAL.py ADDED
@@ -0,0 +1,452 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import gzip
3
+ import torch
4
+ import pathlib
5
+ import requests
6
+ import traceback
7
+ import numpy as np
8
+
9
+ from torch import nn, Tensor
10
+ from torch.nn import functional as F
11
+ from torch.nn.utils.rnn import pack_padded_sequence, pad_sequence
12
+ from torch.distributions import Normal, Categorical
13
+ from typing import *
14
+ from functools import partial
15
+ from itertools import permutations
16
+ try:
17
+ from libriichi3p.mjai import Bot
18
+ from libriichi3p.consts import obs_shape, oracle_obs_shape, ACTION_SPACE, GRP_SIZE
19
+ except:
20
+ import importlib.util
21
+ import sys
22
+ import os
23
+
24
+ # ⚠️ 这里必须填入你在 Colab 中的绝对路径!
25
+ # 假设你的文件在云盘的 MahjongTest 文件夹下,名字叫 libriichi3p.so
26
+ # 如果你的文件叫别的名字,或者在别的文件夹,请务必修改这行路径
27
+ SO_FILE_PATH = "/content/drive/MyDrive/MahjongTest/libriichi3p.so"
28
+
29
+ # 1. 检查文件到底存不存在
30
+ if not os.path.exists(SO_FILE_PATH):
31
+ print(f"❌ 致命错误:在路径 {SO_FILE_PATH} 下根本找不到文件!请检查路径拼写。")
32
+ else:
33
+ print(f"✅ 找到文件: {SO_FILE_PATH},正在尝试强行加载...")
34
+
35
+ try:
36
+ # 2. 根据绝对路径创建模块加载规范 (spec)
37
+ # 第一个参数是你想给它起的名字(供 Python 内部识别),第二个参数是文件路径
38
+ spec = importlib.util.spec_from_file_location("libriichi3p", SO_FILE_PATH)
39
+
40
+ # 3. 实例化模块
41
+ libriichi3p_module = importlib.util.module_from_spec(spec)
42
+
43
+ # 4. 注册到系统的模块字典里 (非常重要!这样后续其他文件 import libriichi3p 就能直接用)
44
+ sys.modules["libriichi3p"] = libriichi3p_module
45
+
46
+ # 5. 执行底层代码加载
47
+ spec.loader.exec_module(libriichi3p_module)
48
+
49
+ print("🎉 强行导入成功!现在可以在代码里正常使用了。")
50
+
51
+ except Exception as e:
52
+ print(f"❌ 导入失败,暴露出真实报错: {e}")
53
+ # ========== Online Server =========== #
54
+ OT_REQUEST_TIMEOUT = 2
55
+ ot_settings = {
56
+ "server": "http://example.com",
57
+ "online": False,
58
+ "api_key": "example_api_key",
59
+ }
60
+ is_online = False
61
+
62
+ def online_settings_init():
63
+ global ot_settings
64
+ # Check if the file exists
65
+ if (pathlib.Path(__file__).parent / 'ot_settings.json').exists():
66
+ with open(pathlib.Path(__file__).parent / 'ot_settings.json', 'r') as f:
67
+ ot_settings = json.load(f)
68
+
69
+ online_settings_init()
70
+ # ==================================== #
71
+
72
+ class ChannelAttention(nn.Module):
73
+ def __init__(self, channels, ratio=16, actv_builder=nn.ReLU, bias=True):
74
+ super().__init__()
75
+ self.shared_mlp = nn.Sequential(
76
+ nn.Linear(channels, channels // ratio, bias=bias),
77
+ actv_builder(),
78
+ nn.Linear(channels // ratio, channels, bias=bias),
79
+ )
80
+ if bias:
81
+ for mod in self.modules():
82
+ if isinstance(mod, nn.Linear):
83
+ nn.init.constant_(mod.bias, 0)
84
+
85
+ def forward(self, x: Tensor):
86
+ avg_out = self.shared_mlp(x.mean(-1))
87
+ max_out = self.shared_mlp(x.amax(-1))
88
+ weight = (avg_out + max_out).sigmoid()
89
+ x = weight.unsqueeze(-1) * x
90
+ return x
91
+
92
+ class ResBlock(nn.Module):
93
+ def __init__(
94
+ self,
95
+ channels,
96
+ *,
97
+ norm_builder = nn.Identity,
98
+ actv_builder = nn.ReLU,
99
+ pre_actv = False,
100
+ ):
101
+ super().__init__()
102
+ self.pre_actv = pre_actv
103
+
104
+ if pre_actv:
105
+ self.res_unit = nn.Sequential(
106
+ norm_builder(),
107
+ actv_builder(),
108
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
109
+ norm_builder(),
110
+ actv_builder(),
111
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
112
+ )
113
+ else:
114
+ self.res_unit = nn.Sequential(
115
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
116
+ norm_builder(),
117
+ actv_builder(),
118
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
119
+ norm_builder(),
120
+ )
121
+ self.actv = actv_builder()
122
+ self.ca = ChannelAttention(channels, actv_builder=actv_builder, bias=True)
123
+
124
+ def forward(self, x):
125
+ out = self.res_unit(x)
126
+ out = self.ca(out)
127
+ out = out + x
128
+ if not self.pre_actv:
129
+ out = self.actv(out)
130
+ return out
131
+
132
+ class ResNet(nn.Module):
133
+ def __init__(
134
+ self,
135
+ in_channels,
136
+ conv_channels,
137
+ num_blocks,
138
+ *,
139
+ norm_builder = nn.Identity,
140
+ actv_builder = nn.ReLU,
141
+ pre_actv = False,
142
+ ):
143
+ super().__init__()
144
+
145
+ blocks = []
146
+ for _ in range(num_blocks):
147
+ blocks.append(ResBlock(
148
+ conv_channels,
149
+ norm_builder = norm_builder,
150
+ actv_builder = actv_builder,
151
+ pre_actv = pre_actv,
152
+ ))
153
+
154
+ layers = [nn.Conv1d(in_channels, conv_channels, kernel_size=3, padding=1, bias=False)]
155
+ if pre_actv:
156
+ layers += [*blocks, norm_builder(), actv_builder()]
157
+ else:
158
+ layers += [norm_builder(), actv_builder(), *blocks]
159
+ layers += [
160
+ nn.Conv1d(conv_channels, 32, kernel_size=3, padding=1),
161
+ actv_builder(),
162
+ nn.Flatten(),
163
+ nn.Linear(32 * 34, 1024),
164
+ ]
165
+ self.net = nn.Sequential(*layers)
166
+
167
+ def forward(self, x):
168
+ return self.net(x)
169
+
170
+ class Brain(nn.Module):
171
+ def __init__(self, *, conv_channels, num_blocks, is_oracle=False, version=1):
172
+ super().__init__()
173
+ self.is_oracle = is_oracle
174
+ self.version = version
175
+
176
+ in_channels = obs_shape(version)[0]
177
+ if is_oracle:
178
+ in_channels += oracle_obs_shape(version)[0]
179
+
180
+ norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01)
181
+ actv_builder = partial(nn.Mish, inplace=True)
182
+ pre_actv = True
183
+
184
+ match version:
185
+ case 1:
186
+ actv_builder = partial(nn.ReLU, inplace=True)
187
+ pre_actv = False
188
+ self.latent_net = nn.Sequential(
189
+ nn.Linear(1024, 512),
190
+ nn.ReLU(inplace=True),
191
+ )
192
+ self.mu_head = nn.Linear(512, 512)
193
+ self.logsig_head = nn.Linear(512, 512)
194
+ case 2:
195
+ pass
196
+ case 3 | 4:
197
+ norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01, eps=1e-3)
198
+ case _:
199
+ raise ValueError(f'Unexpected version {self.version}')
200
+
201
+ self.encoder = ResNet(
202
+ in_channels = in_channels,
203
+ conv_channels = conv_channels,
204
+ num_blocks = num_blocks,
205
+ norm_builder = norm_builder,
206
+ actv_builder = actv_builder,
207
+ pre_actv = pre_actv,
208
+ )
209
+ self.actv = actv_builder()
210
+
211
+ # always use EMA or CMA when True
212
+ self._freeze_bn = False
213
+
214
+ def forward(self, obs: Tensor, invisible_obs: Optional[Tensor] = None) -> Union[Tuple[Tensor, Tensor], Tensor]:
215
+ if self.is_oracle:
216
+ assert invisible_obs is not None
217
+ obs = torch.cat((obs, invisible_obs), dim=1)
218
+ phi = self.encoder(obs)
219
+ phi = F.dropout(phi, p=0.1, training=self.training)
220
+ match self.version:
221
+ case 1:
222
+ latent_out = self.latent_net(phi)
223
+ mu = self.mu_head(latent_out)
224
+ logsig = self.logsig_head(latent_out)
225
+ return mu, logsig
226
+ case 2 | 3 | 4:
227
+ return self.actv(phi)
228
+ case _:
229
+ raise ValueError(f'Unexpected version {self.version}')
230
+
231
+ def train(self, mode=True):
232
+ super().train(mode)
233
+ if self._freeze_bn:
234
+ for mod in self.modules():
235
+ if isinstance(mod, nn.BatchNorm1d):
236
+ mod.eval()
237
+ # I don't think this benefits
238
+ # module.requires_grad_(False)
239
+ return self
240
+
241
+ def reset_running_stats(self):
242
+ for mod in self.modules():
243
+ if isinstance(mod, nn.BatchNorm1d):
244
+ mod.reset_running_stats()
245
+
246
+ def freeze_bn(self, value: bool):
247
+ self._freeze_bn = value
248
+ return self.train(self.training)
249
+
250
+ class AuxNet(nn.Module):
251
+ def __init__(self, dims=None):
252
+ super().__init__()
253
+ self.dims = dims
254
+ self.net = nn.Linear(1024, sum(dims), bias=False)
255
+
256
+ def forward(self, x):
257
+ return self.net(x).split(self.dims, dim=-1)
258
+
259
+ class DQN(nn.Module):
260
+ def __init__(self, *, version=1):
261
+ super().__init__()
262
+ self.version = version
263
+ match version:
264
+ case 1:
265
+ self.v_head = nn.Linear(512, 1)
266
+ self.a_head = nn.Linear(512, ACTION_SPACE)
267
+ case 2 | 3:
268
+ hidden_size = 512 if version == 2 else 256
269
+ self.v_head = nn.Sequential(
270
+ nn.Linear(1024, hidden_size),
271
+ nn.Mish(inplace=True),
272
+ nn.Linear(hidden_size, 1),
273
+ )
274
+ self.a_head = nn.Sequential(
275
+ nn.Linear(1024, hidden_size),
276
+ nn.Mish(inplace=True),
277
+ nn.Linear(hidden_size, ACTION_SPACE),
278
+ )
279
+ case 4:
280
+ self.net = nn.Linear(1024, 1 + ACTION_SPACE)
281
+ nn.init.constant_(self.net.bias, 0)
282
+
283
+ def forward(self, phi, mask):
284
+ if self.version == 4:
285
+ v, a = self.net(phi).split((1, ACTION_SPACE), dim=-1)
286
+ else:
287
+ v = self.v_head(phi)
288
+ a = self.a_head(phi)
289
+ a_sum = a.masked_fill(~mask, 0.).sum(-1, keepdim=True)
290
+ mask_sum = mask.sum(-1, keepdim=True)
291
+ a_mean = a_sum / mask_sum
292
+ q = (v + a - a_mean).masked_fill(~mask, -1e9)
293
+ return q
294
+
295
+
296
+ class MortalEngine:
297
+ def __init__(
298
+ self,
299
+ brain,
300
+ dqn,
301
+ is_oracle,
302
+ version,
303
+ device = None,
304
+ stochastic_latent = False,
305
+ enable_amp = False,
306
+ enable_quick_eval = True,
307
+ enable_rule_based_agari_guard = False,
308
+ name = 'NoName',
309
+ boltzmann_epsilon = 0,
310
+ boltzmann_temp = 1,
311
+ top_p = 1,
312
+ ):
313
+ self.engine_type = 'mortal'
314
+ self.device = device or torch.device('cpu')
315
+ assert isinstance(self.device, torch.device)
316
+ self.brain = brain.to(self.device).eval()
317
+ self.dqn = dqn.to(self.device).eval()
318
+ self.is_oracle = is_oracle
319
+ self.version = version
320
+ self.stochastic_latent = stochastic_latent
321
+
322
+ self.enable_amp = enable_amp
323
+ self.enable_quick_eval = enable_quick_eval
324
+ self.enable_rule_based_agari_guard = enable_rule_based_agari_guard
325
+ self.name = name
326
+
327
+ self.boltzmann_epsilon = boltzmann_epsilon
328
+ self.boltzmann_temp = boltzmann_temp
329
+ self.top_p = top_p
330
+
331
+ def react_batch(self, obs, masks, invisible_obs):
332
+ # ========== Online Server =========== #
333
+ global ot_settings, is_online
334
+ # print('Reacting Batch')
335
+ if ot_settings['online']:
336
+ try:
337
+ list_obs = [o.tolist() for o in obs]
338
+ list_masks = [m.tolist() for m in masks]
339
+ post_data = {
340
+ 'obs': list_obs,
341
+ 'masks': list_masks,
342
+ }
343
+ data = json.dumps(post_data, separators=(',', ':'))
344
+ compressed_data = gzip.compress(data.encode('utf-8'))
345
+ headers = {
346
+ 'Authorization': ot_settings['api_key'],
347
+ 'Content-Encoding': 'gzip',
348
+ }
349
+ r = requests.post(
350
+ f'{ot_settings["server"]}/react_batch_3p',
351
+ headers=headers,
352
+ data=compressed_data,
353
+ timeout=OT_REQUEST_TIMEOUT
354
+ )
355
+ assert r.status_code == 200
356
+ is_online = True
357
+ r_json = r.json()
358
+ return r_json['actions'], r_json['q_out'], r_json['masks'], r_json['is_greedy']
359
+ except:
360
+ is_online = False
361
+ pass
362
+ # ==================================== #
363
+ try:
364
+ with (
365
+ torch.autocast(self.device.type, enabled=self.enable_amp),
366
+ torch.inference_mode(),
367
+ ):
368
+ return self._react_batch(obs, masks, invisible_obs)
369
+ except Exception as ex:
370
+ raise Exception(f'{ex}\n{traceback.format_exc()}')
371
+
372
+ def _react_batch(self, obs, masks, invisible_obs):
373
+ obs = torch.as_tensor(np.stack(obs, axis=0), device=self.device)
374
+ masks = torch.as_tensor(np.stack(masks, axis=0), device=self.device)
375
+ invisible_obs = None
376
+ if self.is_oracle:
377
+ invisible_obs = torch.as_tensor(np.stack(invisible_obs, axis=0), device=self.device)
378
+ batch_size = obs.shape[0]
379
+
380
+ match self.version:
381
+ case 1:
382
+ mu, logsig = self.brain(obs, invisible_obs)
383
+ if self.stochastic_latent:
384
+ latent = Normal(mu, logsig.exp() + 1e-6).sample()
385
+ else:
386
+ latent = mu
387
+ q_out = self.dqn(latent, masks)
388
+ case 2 | 3 | 4:
389
+ phi = self.brain(obs)
390
+ q_out = self.dqn(phi, masks)
391
+
392
+ if self.boltzmann_epsilon > 0:
393
+ is_greedy = torch.full((batch_size,), 1-self.boltzmann_epsilon, device=self.device).bernoulli().to(torch.bool)
394
+ logits = (q_out / self.boltzmann_temp).masked_fill(~masks, -torch.inf)
395
+ sampled = sample_top_p(logits, self.top_p)
396
+ actions = torch.where(is_greedy, q_out.argmax(-1), sampled)
397
+ else:
398
+ is_greedy = torch.ones(batch_size, dtype=torch.bool, device=self.device)
399
+ actions = q_out.argmax(-1)
400
+ return actions.tolist(), q_out.tolist(), masks.tolist(), is_greedy.tolist()
401
+
402
+ def sample_top_p(logits, p):
403
+ if p >= 1:
404
+ return Categorical(logits=logits).sample()
405
+ if p <= 0:
406
+ return logits.argmax(-1)
407
+ probs = logits.softmax(-1)
408
+ probs_sort, probs_idx = probs.sort(-1, descending=True)
409
+ probs_sum = probs_sort.cumsum(-1)
410
+ mask = probs_sum - probs_sort > p
411
+ probs_sort[mask] = 0.
412
+ sampled = probs_idx.gather(-1, probs_sort.multinomial(1)).squeeze(-1)
413
+ return sampled
414
+
415
+ def load_model(seat: int, model: str) -> Bot:
416
+ # check if GPU is available
417
+ if torch.cuda.is_available():
418
+ device = torch.device('cuda')
419
+ else:
420
+ device = torch.device('cpu')
421
+
422
+ # latest binary model
423
+ if model == None:
424
+ model = 'Elite4zWeightedBest5.pth'
425
+ model = str(model).split('?')[0]
426
+ control_state_file = model
427
+ print(control_state_file, 'loaded')
428
+
429
+ # Get the path of control_state_file = current directory / control_state_file
430
+ control_state_file = pathlib.Path(__file__).parent / control_state_file
431
+ state = torch.load(control_state_file, map_location=device)
432
+
433
+ mortal = Brain(version=state['config']['control']['version'], conv_channels=state['config']['resnet']['conv_channels'], num_blocks=state['config']['resnet']['num_blocks']).eval()
434
+ dqn = DQN(version=state['config']['control']['version']).eval()
435
+ mortal.load_state_dict(state['mortal'])
436
+ dqn.load_state_dict(state['current_dqn'])
437
+
438
+ engine = MortalEngine(
439
+ mortal,
440
+ dqn,
441
+ is_oracle = False,
442
+ version = state['config']['control']['version'],
443
+ device = device,
444
+ enable_amp = False,
445
+ enable_quick_eval = False,
446
+ enable_rule_based_agari_guard = True,
447
+ name = 'mortal',
448
+ top_p = 1,
449
+ )
450
+
451
+ bot = Bot(engine, seat)
452
+ return bot
model3pNEW.py ADDED
@@ -0,0 +1,418 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import gzip
3
+ import torch
4
+ import pathlib
5
+ import requests
6
+ import traceback
7
+ import numpy as np
8
+
9
+ from torch import nn, Tensor
10
+ from torch.nn import functional as F
11
+ from torch.nn.utils.rnn import pack_padded_sequence, pad_sequence
12
+ from torch.distributions import Normal, Categorical
13
+ from typing import *
14
+ from functools import partial
15
+ from itertools import permutations
16
+ from libriichiSanma.mjai import Bot
17
+ from libriichiSanma.consts import obs_shape, oracle_obs_shape, ACTION_SPACE, GRP_SIZE
18
+
19
+ # ========== Online Server =========== #
20
+ OT_REQUEST_TIMEOUT = 2
21
+ ot_settings = {
22
+ "server": "http://example.com",
23
+ "online": False,
24
+ "api_key": "example_api_key",
25
+ }
26
+ is_online = False
27
+
28
+ def online_settings_init():
29
+ global ot_settings
30
+ # Check if the file exists
31
+ if (pathlib.Path(__file__).parent / 'ot_settings.json').exists():
32
+ with open(pathlib.Path(__file__).parent / 'ot_settings.json', 'r') as f:
33
+ ot_settings = json.load(f)
34
+
35
+ online_settings_init()
36
+ # ==================================== #
37
+
38
+ class ChannelAttention(nn.Module):
39
+ def __init__(self, channels, ratio=16, actv_builder=nn.ReLU, bias=True):
40
+ super().__init__()
41
+ self.shared_mlp = nn.Sequential(
42
+ nn.Linear(channels, channels // ratio, bias=bias),
43
+ actv_builder(),
44
+ nn.Linear(channels // ratio, channels, bias=bias),
45
+ )
46
+ if bias:
47
+ for mod in self.modules():
48
+ if isinstance(mod, nn.Linear):
49
+ nn.init.constant_(mod.bias, 0)
50
+
51
+ def forward(self, x: Tensor):
52
+ avg_out = self.shared_mlp(x.mean(-1))
53
+ max_out = self.shared_mlp(x.amax(-1))
54
+ weight = (avg_out + max_out).sigmoid()
55
+ x = weight.unsqueeze(-1) * x
56
+ return x
57
+
58
+ class ResBlock(nn.Module):
59
+ def __init__(
60
+ self,
61
+ channels,
62
+ *,
63
+ norm_builder = nn.Identity,
64
+ actv_builder = nn.ReLU,
65
+ pre_actv = False,
66
+ ):
67
+ super().__init__()
68
+ self.pre_actv = pre_actv
69
+
70
+ if pre_actv:
71
+ self.res_unit = nn.Sequential(
72
+ norm_builder(),
73
+ actv_builder(),
74
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
75
+ norm_builder(),
76
+ actv_builder(),
77
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
78
+ )
79
+ else:
80
+ self.res_unit = nn.Sequential(
81
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
82
+ norm_builder(),
83
+ actv_builder(),
84
+ nn.Conv1d(channels, channels, kernel_size=3, padding=1, bias=False),
85
+ norm_builder(),
86
+ )
87
+ self.actv = actv_builder()
88
+ self.ca = ChannelAttention(channels, actv_builder=actv_builder, bias=True)
89
+
90
+ def forward(self, x):
91
+ out = self.res_unit(x)
92
+ out = self.ca(out)
93
+ out = out + x
94
+ if not self.pre_actv:
95
+ out = self.actv(out)
96
+ return out
97
+
98
+ class ResNet(nn.Module):
99
+ def __init__(
100
+ self,
101
+ in_channels,
102
+ conv_channels,
103
+ num_blocks,
104
+ *,
105
+ norm_builder = nn.Identity,
106
+ actv_builder = nn.ReLU,
107
+ pre_actv = False,
108
+ ):
109
+ super().__init__()
110
+
111
+ blocks = []
112
+ for _ in range(num_blocks):
113
+ blocks.append(ResBlock(
114
+ conv_channels,
115
+ norm_builder = norm_builder,
116
+ actv_builder = actv_builder,
117
+ pre_actv = pre_actv,
118
+ ))
119
+
120
+ layers = [nn.Conv1d(in_channels, conv_channels, kernel_size=3, padding=1, bias=False)]
121
+ if pre_actv:
122
+ layers += [*blocks, norm_builder(), actv_builder()]
123
+ else:
124
+ layers += [norm_builder(), actv_builder(), *blocks]
125
+ layers += [
126
+ nn.Conv1d(conv_channels, 32, kernel_size=3, padding=1),
127
+ actv_builder(),
128
+ nn.Dropout(p=0.05),
129
+ nn.Flatten(),
130
+ nn.Linear(32 * 34, 1024),
131
+ ]
132
+ self.net = nn.Sequential(*layers)
133
+
134
+ def forward(self, x):
135
+ return self.net(x)
136
+
137
+ class Brain(nn.Module):
138
+ def __init__(self, *, conv_channels, num_blocks, is_oracle=False, version=1):
139
+ super().__init__()
140
+ self.is_oracle = is_oracle
141
+ self.version = version
142
+
143
+ in_channels = obs_shape(version)[0]
144
+ if is_oracle:
145
+ in_channels += oracle_obs_shape(version)[0]
146
+
147
+ norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01)
148
+ actv_builder = partial(nn.Mish, inplace=True)
149
+ pre_actv = True
150
+
151
+ match version:
152
+ case 1:
153
+ actv_builder = partial(nn.ReLU, inplace=True)
154
+ pre_actv = False
155
+ self.latent_net = nn.Sequential(
156
+ nn.Linear(1024, 512),
157
+ nn.ReLU(inplace=True),
158
+ )
159
+ self.mu_head = nn.Linear(512, 512)
160
+ self.logsig_head = nn.Linear(512, 512)
161
+ case 2:
162
+ pass
163
+ case 3 | 4:
164
+ norm_builder = partial(nn.BatchNorm1d, conv_channels, momentum=0.01, eps=1e-3)
165
+ case _:
166
+ raise ValueError(f'Unexpected version {self.version}')
167
+
168
+ self.encoder = ResNet(
169
+ in_channels = in_channels,
170
+ conv_channels = conv_channels,
171
+ num_blocks = num_blocks,
172
+ norm_builder = norm_builder,
173
+ actv_builder = actv_builder,
174
+ pre_actv = pre_actv,
175
+ )
176
+ self.actv = actv_builder()
177
+
178
+ # always use EMA or CMA when True
179
+ self._freeze_bn = False
180
+
181
+ def forward(self, obs: Tensor, invisible_obs: Optional[Tensor] = None) -> Union[Tuple[Tensor, Tensor], Tensor]:
182
+ if self.is_oracle:
183
+ assert invisible_obs is not None
184
+ obs = torch.cat((obs, invisible_obs), dim=1)
185
+ phi = self.encoder(obs)
186
+ phi = F.dropout(phi, p=0.1, training=self.training)
187
+ match self.version:
188
+ case 1:
189
+ latent_out = self.latent_net(phi)
190
+ mu = self.mu_head(latent_out)
191
+ logsig = self.logsig_head(latent_out)
192
+ return mu, logsig
193
+ case 2 | 3 | 4:
194
+ return self.actv(phi)
195
+ case _:
196
+ raise ValueError(f'Unexpected version {self.version}')
197
+
198
+ def train(self, mode=True):
199
+ super().train(mode)
200
+ if self._freeze_bn:
201
+ for mod in self.modules():
202
+ if isinstance(mod, nn.BatchNorm1d):
203
+ mod.eval()
204
+ # I don't think this benefits
205
+ # module.requires_grad_(False)
206
+ return self
207
+
208
+ def reset_running_stats(self):
209
+ for mod in self.modules():
210
+ if isinstance(mod, nn.BatchNorm1d):
211
+ mod.reset_running_stats()
212
+
213
+ def freeze_bn(self, value: bool):
214
+ self._freeze_bn = value
215
+ return self.train(self.training)
216
+
217
+ class AuxNet(nn.Module):
218
+ def __init__(self, dims=None):
219
+ super().__init__()
220
+ self.dims = dims
221
+ self.net = nn.Linear(1024, sum(dims), bias=False)
222
+
223
+ def forward(self, x):
224
+ return self.net(x).split(self.dims, dim=-1)
225
+
226
+ class DQN(nn.Module):
227
+ def __init__(self, *, version=1):
228
+ super().__init__()
229
+ self.version = version
230
+ match version:
231
+ case 1:
232
+ self.v_head = nn.Linear(512, 1)
233
+ self.a_head = nn.Linear(512, ACTION_SPACE)
234
+ case 2 | 3:
235
+ hidden_size = 512 if version == 2 else 256
236
+ self.v_head = nn.Sequential(
237
+ nn.Linear(1024, hidden_size),
238
+ nn.Mish(inplace=True),
239
+ nn.Linear(hidden_size, 1),
240
+ )
241
+ self.a_head = nn.Sequential(
242
+ nn.Linear(1024, hidden_size),
243
+ nn.Mish(inplace=True),
244
+ nn.Linear(hidden_size, ACTION_SPACE),
245
+ )
246
+ case 4:
247
+ self.net = nn.Linear(1024, 1 + ACTION_SPACE)
248
+ nn.init.constant_(self.net.bias, 0)
249
+
250
+ def forward(self, phi, mask):
251
+ if self.version == 4:
252
+ v, a = self.net(phi).split((1, ACTION_SPACE), dim=-1)
253
+ else:
254
+ v = self.v_head(phi)
255
+ a = self.a_head(phi)
256
+ a_sum = a.masked_fill(~mask, 0.).sum(-1, keepdim=True)
257
+ mask_sum = mask.sum(-1, keepdim=True)
258
+ a_mean = a_sum / mask_sum
259
+ q = (v + a - a_mean).masked_fill(~mask, -1e9)
260
+ return q
261
+
262
+
263
+ class MortalEngine:
264
+ def __init__(
265
+ self,
266
+ brain,
267
+ dqn,
268
+ is_oracle,
269
+ version,
270
+ device = None,
271
+ stochastic_latent = False,
272
+ enable_amp = False,
273
+ enable_quick_eval = True,
274
+ enable_rule_based_agari_guard = False,
275
+ name = 'NoName',
276
+ boltzmann_epsilon = 0,
277
+ boltzmann_temp = 1,
278
+ top_p = 1,
279
+ ):
280
+ self.engine_type = 'mortal'
281
+ self.device = device or torch.device('cpu')
282
+ assert isinstance(self.device, torch.device)
283
+ self.brain = brain.to(self.device).eval()
284
+ self.dqn = dqn.to(self.device).eval()
285
+ self.is_oracle = is_oracle
286
+ self.version = version
287
+ self.stochastic_latent = stochastic_latent
288
+
289
+ self.enable_amp = enable_amp
290
+ self.enable_quick_eval = enable_quick_eval
291
+ self.enable_rule_based_agari_guard = enable_rule_based_agari_guard
292
+ self.name = name
293
+
294
+ self.boltzmann_epsilon = boltzmann_epsilon
295
+ self.boltzmann_temp = boltzmann_temp
296
+ self.top_p = top_p
297
+
298
+ def react_batch(self, obs, masks, invisible_obs):
299
+ # ========== Online Server =========== #
300
+ global ot_settings, is_online
301
+ # print('Reacting Batch')
302
+ if ot_settings['online']:
303
+ try:
304
+ list_obs = [o.tolist() for o in obs]
305
+ list_masks = [m.tolist() for m in masks]
306
+ post_data = {
307
+ 'obs': list_obs,
308
+ 'masks': list_masks,
309
+ }
310
+ data = json.dumps(post_data, separators=(',', ':'))
311
+ compressed_data = gzip.compress(data.encode('utf-8'))
312
+ headers = {
313
+ 'Authorization': ot_settings['api_key'],
314
+ 'Content-Encoding': 'gzip',
315
+ }
316
+ r = requests.post(
317
+ f'{ot_settings["server"]}/react_batch_3p',
318
+ headers=headers,
319
+ data=compressed_data,
320
+ timeout=OT_REQUEST_TIMEOUT
321
+ )
322
+ assert r.status_code == 200
323
+ is_online = True
324
+ r_json = r.json()
325
+ return r_json['actions'], r_json['q_out'], r_json['masks'], r_json['is_greedy']
326
+ except:
327
+ is_online = False
328
+ pass
329
+ # ==================================== #
330
+ try:
331
+ with (
332
+ torch.autocast(self.device.type, enabled=self.enable_amp),
333
+ torch.inference_mode(),
334
+ ):
335
+ return self._react_batch(obs, masks, invisible_obs)
336
+ except Exception as ex:
337
+ raise Exception(f'{ex}\n{traceback.format_exc()}')
338
+
339
+ def _react_batch(self, obs, masks, invisible_obs):
340
+ obs = torch.as_tensor(np.stack(obs, axis=0), device=self.device)
341
+ masks = torch.as_tensor(np.stack(masks, axis=0), device=self.device)
342
+ invisible_obs = None
343
+ if self.is_oracle:
344
+ invisible_obs = torch.as_tensor(np.stack(invisible_obs, axis=0), device=self.device)
345
+ batch_size = obs.shape[0]
346
+
347
+ match self.version:
348
+ case 1:
349
+ mu, logsig = self.brain(obs, invisible_obs)
350
+ if self.stochastic_latent:
351
+ latent = Normal(mu, logsig.exp() + 1e-6).sample()
352
+ else:
353
+ latent = mu
354
+ q_out = self.dqn(latent, masks)
355
+ case 2 | 3 | 4:
356
+ phi = self.brain(obs)
357
+ q_out = self.dqn(phi, masks)
358
+
359
+ if self.boltzmann_epsilon > 0:
360
+ is_greedy = torch.full((batch_size,), 1-self.boltzmann_epsilon, device=self.device).bernoulli().to(torch.bool)
361
+ logits = (q_out / self.boltzmann_temp).masked_fill(~masks, -torch.inf)
362
+ sampled = sample_top_p(logits, self.top_p)
363
+ actions = torch.where(is_greedy, q_out.argmax(-1), sampled)
364
+ else:
365
+ is_greedy = torch.ones(batch_size, dtype=torch.bool, device=self.device)
366
+ actions = q_out.argmax(-1)
367
+ return actions.tolist(), q_out.tolist(), masks.tolist(), is_greedy.tolist()
368
+
369
+ def sample_top_p(logits, p):
370
+ if p >= 1:
371
+ return Categorical(logits=logits).sample()
372
+ if p <= 0:
373
+ return logits.argmax(-1)
374
+ probs = logits.softmax(-1)
375
+ probs_sort, probs_idx = probs.sort(-1, descending=True)
376
+ probs_sum = probs_sort.cumsum(-1)
377
+ mask = probs_sum - probs_sort > p
378
+ probs_sort[mask] = 0.
379
+ sampled = probs_idx.gather(-1, probs_sort.multinomial(1)).squeeze(-1)
380
+ return sampled
381
+
382
+ def load_model(seat: int, model: str) -> Bot:
383
+ # check if GPU is available
384
+ if torch.cuda.is_available() or False:
385
+ device = torch.device('cuda')
386
+ else:
387
+ device = torch.device('cpu')
388
+
389
+ # latest binary model
390
+ control_state_file = model
391
+ print(control_state_file, 'loaded')
392
+
393
+ # Get the path of control_state_file = current directory / control_state_file
394
+ control_state_file = pathlib.Path(__file__).parent / control_state_file
395
+ state = torch.load(control_state_file, map_location=device)
396
+
397
+ mortal = Brain(version=state['config']['control']['version'], conv_channels=state['config']['resnet']['conv_channels'], num_blocks=state['config']['resnet']['num_blocks']).eval()
398
+ dqn = DQN(version=state['config']['control']['version']).eval()
399
+ mortal_key = 'student_brain' if 'student_brain' in state else 'mortal'
400
+ dqn_key = 'student_dqn' if 'student_dqn' in state else 'current_dqn'
401
+ mortal.load_state_dict(state[mortal_key])
402
+ dqn.load_state_dict(state[dqn_key])
403
+
404
+ engine = MortalEngine(
405
+ mortal,
406
+ dqn,
407
+ is_oracle = False,
408
+ version = state['config']['control']['version'],
409
+ device = device,
410
+ enable_amp = False,
411
+ enable_quick_eval = False,
412
+ enable_rule_based_agari_guard = True,
413
+ name = 'mortal',
414
+ top_p = 1,
415
+ )
416
+
417
+ bot = Bot(engine, seat)
418
+ return bot