ffzeroHua commited on
Commit
e0fab04
·
verified ·
1 Parent(s): f598941

Upload 8 files

Browse files
.gitattributes CHANGED
@@ -33,3 +33,5 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ libriichi.so filter=lfs diff=lfs merge=lfs -text
37
+ libriichi3p.so filter=lfs diff=lfs merge=lfs -text
Backup_B50_Rank1.75_Rwd0.12_Ent0.149.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:eda491eb91d98186122777dd93b89b9786be05ee90d271506f416e661f68df6d
3
+ size 273678105
Bot.py ADDED
@@ -0,0 +1,235 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import sys
3
+ import model
4
+ import model3p
5
+ import numpy as np
6
+
7
+ tiles_tenhou = {
8
+ '1m': 0, '2m': 1, '3m': 2, '4m': 3, '5m': 4, '5mr': 4.5, '6m': 5, '7m': 6, '8m': 7, '9m': 8,
9
+ '1p': 9, '2p': 10, '3p': 11, '4p': 12, '5p': 13, '5pr': 13.5, '6p': 14, '7p': 15, '8p': 16, '9p': 17,
10
+ '1s': 18, '2s': 19, '3s': 20, '4s': 21, '5s': 22, '5sr': 22.5, '6s': 23, '7s': 24, '8s': 25, '9s': 26,
11
+ 'E': 27, 'S': 28, 'W': 29, 'N': 30, 'P': 31, 'F': 32, 'C': 33
12
+ }
13
+
14
+ MASK_4P = [
15
+ "1m", "2m", "3m", "4m", "5m", "6m", "7m", "8m", "9m",
16
+ "1p", "2p", "3p", "4p", "5p", "6p", "7p", "8p", "9p",
17
+ "1s", "2s", "3s", "4s", "5s", "6s", "7s", "8s", "9s",
18
+ "E", "S", "W", "N", "P", "F", "C",
19
+ '5mr', '5pr', '5sr',
20
+ 'reach', 'chi_low', 'chi_mid', 'chi_high', 'pon', 'kan', 'hora', 'ryukyoku', 'none'
21
+ ]
22
+
23
+ MASK_3P = [
24
+ "1m", "2m", "3m", "4m", "5m", "6m", "7m", "8m", "9m",
25
+ "1p", "2p", "3p", "4p", "5p", "6p", "7p", "8p", "9p",
26
+ "1s", "2s", "3s", "4s", "5s", "6s", "7s", "8s", "9s",
27
+ "E", "S", "W", "N", "P", "F", "C",
28
+ '5mr', '5pr', '5sr',
29
+ 'reach', 'pon', 'kan', 'nukidora', 'hora', 'ryukyoku', 'none'
30
+ ]
31
+
32
+ def SoftMax(arr, temperature=1.0):
33
+ arr = np.array(arr, dtype=float) # Ensure the input is a numpy array of floats
34
+ if arr.size == 0:
35
+ return arr # Return the empty array if input is empty
36
+ if not temperature == 1.0:
37
+ arr /= temperature # Scale by temperature if temperature is not approximately 1
38
+ # Shift values by max for numerical stability
39
+ max_val = np.max(arr)
40
+ arr = arr - max_val
41
+ # Apply the softmax transformation
42
+ exp_arr = np.exp(arr)
43
+ sum_exp = np.sum(exp_arr)
44
+ softmax_arr = exp_arr / sum_exp
45
+ return softmax_arr
46
+
47
+ def ToBinStr(mask_bits):
48
+ binary_string = bin(mask_bits)[2:]
49
+ binary_string = binary_string.zfill(46)
50
+ return binary_string
51
+
52
+
53
+ def ToBoolList(mask_bits):
54
+ binary_string = ToBinStr(mask_bits)
55
+ bool_list = []
56
+ for bit in binary_string[::-1]:
57
+ bool_list.append(bit == '1')
58
+ return bool_list
59
+
60
+ def ParseMeta(is_3p, meta):
61
+ if is_3p:
62
+ mask_list = MASK_3P
63
+ else:
64
+ mask_list = MASK_4P
65
+
66
+ q_values = meta['q_values']
67
+ mask_bits = meta['mask_bits']
68
+ mask = ToBoolList(mask_bits)
69
+ weight_values = SoftMax(q_values)
70
+
71
+ q_value_idx = 0
72
+ option_list = []
73
+ for i in range(46):
74
+ if mask[i]:
75
+ option_list.append((mask_list[i], weight_values[q_value_idx]))
76
+ q_value_idx += 1
77
+
78
+ option_list = sorted(option_list, key=lambda x: x[1], reverse=True)
79
+ return option_list
80
+
81
+
82
+ class Bot:
83
+ def __init__(self):
84
+ self.player_id: int = None
85
+ self.model = None
86
+
87
+ def react(self, events: str):
88
+ events = json.loads(events)
89
+ return_action = None
90
+ for e in events:
91
+ if e["type"] == "start_game":
92
+ self.player_id = e["id"]
93
+ model_type = e.get('model')
94
+ self.model = model3p.load_model(self.player_id, model_type)
95
+ continue
96
+ if self.model is None or self.player_id is None:
97
+ raise Exception(f"Model is not loaded yet")
98
+ continue
99
+ if e["type"] == "end_game":
100
+ self.player_id = None
101
+ self.model = None
102
+ continue
103
+ return_action = self.model.react(json.dumps(e, separators=(",", ":")))
104
+ return return_action
105
+
106
+ class Bot4P:
107
+ def __init__(self):
108
+ self.player_id: int = None
109
+ self.model = None
110
+
111
+ def react(self, events: str, H = True):
112
+ events = json.loads(events)
113
+ return_action = None
114
+ for e in events:
115
+ if e["type"] == "start_game":
116
+ self.player_id = e["id"]
117
+ model_type = e.get('model')
118
+ self.model = model.load_model(self.player_id, model_type)
119
+ continue
120
+ if self.model is None or self.player_id is None:
121
+ raise Exception(f"Model is not loaded yet")
122
+ continue
123
+ if e["type"] == "end_game":
124
+ self.player_id = None
125
+ self.model = None
126
+ continue
127
+ return_action = self.model.react(json.dumps(e, separators=(",", ":")))
128
+ if H:
129
+ return return_action
130
+ if return_action == None:
131
+ return None
132
+ original = json.loads(return_action)
133
+ return Select(original, ParseMeta(False, original['meta']))
134
+ false = False
135
+ true = True
136
+ if __name__ == '__main__':
137
+ bot = Bot4P()
138
+ print(bot)
139
+
140
+ events = {
141
+ "is3p": false,
142
+ "model": "v2-a",
143
+ "events": [
144
+ {
145
+ "id": 0,
146
+ "type": "start_game"
147
+ },
148
+ {
149
+ "oya": 0,
150
+ "type": "start_kyoku",
151
+ "honba": 0,
152
+ "kyoku": 1,
153
+ "bakaze": "E",
154
+ "scores": [
155
+ 25000,
156
+ 25000,
157
+ 25000,
158
+ 0
159
+ ],
160
+ "tehais": [
161
+ [
162
+ "1m",
163
+ "2s",
164
+ "3s",
165
+ "3p",
166
+ "4p",
167
+ "5p",
168
+ "5s",
169
+ "6s",
170
+ "7s",
171
+ "N",
172
+ "N",
173
+ "S",
174
+ "W"
175
+ ],
176
+ [
177
+ "?",
178
+ "?",
179
+ "?",
180
+ "?",
181
+ "?",
182
+ "?",
183
+ "?",
184
+ "?",
185
+ "?",
186
+ "?",
187
+ "?",
188
+ "?",
189
+ "?"
190
+ ],
191
+ [
192
+ "?",
193
+ "?",
194
+ "?",
195
+ "?",
196
+ "?",
197
+ "?",
198
+ "?",
199
+ "?",
200
+ "?",
201
+ "?",
202
+ "?",
203
+ "?",
204
+ "?"
205
+ ],
206
+ [
207
+ "?",
208
+ "?",
209
+ "?",
210
+ "?",
211
+ "?",
212
+ "?",
213
+ "?",
214
+ "?",
215
+ "?",
216
+ "?",
217
+ "?",
218
+ "?",
219
+ "?"
220
+ ]
221
+ ],
222
+ "kyotaku": 0,
223
+ "dora_marker": "9m"
224
+ },
225
+ {
226
+ "pai": "1p",
227
+ "type": "tsumo",
228
+ "actor": 0
229
+ }
230
+ ],
231
+ "timestamp": "2025-09-22T15:03:13.873Z"
232
+ }
233
+
234
+ res = bot.react(json.dumps(events['events']))
235
+ print(res)
Elite4z9070.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5d9e6fd281df9eb1150d13d034265c252b967a1e2046cd1f3a2a5b4d56cb198d
3
+ size 273664233
libriichi.so ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3aeb09112198f6babe9bba4dd392ea009669ee7ba5d33cd3cd4242c737bb004b
3
+ size 21165448
libriichi3p.so ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:03900834051021f662fec35c6e9608f4d4c5aa61b4c4ce37b49fa2e861bf619b
3
+ size 1873424
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
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torch
2
+ orjson
3
+ gradio
4
+ matplotlib
5
+ pandas
6
+ riichienv
本地训练器.py ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+
3
+ import orjson
4
+ import concurrent.futures
5
+ import random
6
+ import torch
7
+ from riichienv import RiichiEnv, GameRule
8
+ from model3pLOCAL import load_model
9
+
10
+ # ==========================================
11
+ # 1. 高性能 Patch 逻辑 (保持不变)
12
+ # ==========================================
13
+ def patch_event_fast(event_str):
14
+ if '"kita"' in event_str:
15
+ event_str = event_str.replace('"kita"', '"nukidora"')
16
+
17
+ if '"start_kyoku"' in event_str or '"deltas"' in event_str:
18
+ event = orjson.loads(event_str)
19
+ if event.get('type') == 'start_kyoku':
20
+ scores = event.setdefault('scores', [])
21
+ while len(scores) < 4:
22
+ scores.append(0)
23
+ tehais = event.setdefault('tehais', [])
24
+ while len(tehais) < 4:
25
+ tehais.append(["?" for _ in range(13)])
26
+ if 'deltas' in event:
27
+ deltas = event['deltas']
28
+ while len(deltas) < 4:
29
+ deltas.append(0)
30
+ return orjson.dumps(event).decode('utf-8')
31
+ return event_str
32
+
33
+ def patch_resp_fast(resp_str):
34
+ if not resp_str:
35
+ return resp_str
36
+ return resp_str.replace('"nukidora"', '"kita"')
37
+
38
+ # ==========================================
39
+ # 2. 全局模型缓存
40
+ # ==========================================
41
+ _MODEL_CACHE = {}
42
+
43
+ def get_cached_model(player_id: int, model_file: str):
44
+ key = (player_id, model_file)
45
+ if key not in _MODEL_CACHE:
46
+ _MODEL_CACHE[key] = load_model(player_id, model_file)
47
+ return _MODEL_CACHE[key]
48
+
49
+ class MortalAgent:
50
+ def __init__(self, player_id: int, model_file: str):
51
+ self.player_id = player_id
52
+ self.model = get_cached_model(player_id, model_file)
53
+
54
+ def act(self, obs):
55
+ resp = None
56
+ for event in obs.new_events():
57
+ event_patched = patch_event_fast(event)
58
+ resp = patch_resp_fast(self.model.react(event_patched))
59
+ action = obs.select_action_from_mjai(resp)
60
+ assert action is not None, "Mortal must return a legal action"
61
+ return action
62
+
63
+ # ==========================================
64
+ # 3. 单局游戏任务 (加入随机座次)
65
+ # ==========================================
66
+ def play_one_game(game_index):
67
+ env = RiichiEnv(game_mode="3p-red-half", rule=GameRule.default_tenhou())
68
+
69
+ # --- 随机分配座次 ---
70
+ # 随机选一个位置给 NEW_MODEL (0, 1, 或 2)
71
+ new_seat = random.randrange(3)
72
+
73
+ agents = {}
74
+ for i in range(3):
75
+ model_file = "Backup_B50_Rank1.75_Rwd0.12_Ent0.149.pth" if i == new_seat else "Elite4z9070.pth"
76
+ agents[i] = MortalAgent(i, model_file)
77
+
78
+ obs_dict = env.reset()
79
+ while not env.done():
80
+ actions = {pid: agents[pid].act(obs) for pid, obs in obs_dict.items()}
81
+ obs_dict = env.step(actions)
82
+
83
+ scores = env.scores()
84
+ ranks = env.ranks()
85
+
86
+ # 根据随机到的 new_seat 提取新模型的成绩
87
+ # 顺位:ranks[new_seat], 分数:scores[new_seat]
88
+ return ranks[new_seat], scores[new_seat]
89
+
90
+ # ==========================================
91
+ # 4. 多进程调度主循环
92
+ # ==========================================
93
+ if __name__ == "__main__":
94
+ REPORT_FILE = "训练报告-9070旧-W5新.txt"
95
+ NUM_WORKERS = 1 # 建议设为 CPU 核心数或稍小
96
+
97
+ print(f"🚀 启动评估程序,当前多开数: {NUM_WORKERS}")
98
+ print(f"📝 成绩将随机化座次并追加写入到 {REPORT_FILE}")
99
+ print("⏹️ 随时按 Ctrl+C 即可安全停止\n")
100
+
101
+ try:
102
+ with concurrent.futures.ProcessPoolExecutor(max_workers=NUM_WORKERS) as executor:
103
+ # 维持任务队列
104
+ futures = {executor.submit(play_one_game, i) for i in range(NUM_WORKERS * 2)}
105
+ games_completed = 0
106
+
107
+ while futures:
108
+ done, futures = concurrent.futures.wait(
109
+ futures, return_when=concurrent.futures.FIRST_COMPLETED
110
+ )
111
+
112
+ # 写入汇总
113
+ with open(REPORT_FILE, "a") as f:
114
+ for future in done:
115
+ try:
116
+ rank, score = future.result()
117
+ f.write(f"{rank} {score}\n")
118
+ f.flush()
119
+
120
+ games_completed += 1
121
+ print(f"[{games_completed}局完成] 新模型成绩 -> 顺位: {rank}, 分数: {score}")
122
+
123
+ except Exception as e:
124
+ print(f"⚠️ 异常: {e}")
125
+
126
+ # 补齐新任务
127
+ futures.add(executor.submit(play_one_game, games_completed))
128
+
129
+ except KeyboardInterrupt:
130
+ print("\n🛑 评估已手动停止。")