zhoudoe23 commited on
Commit
f672d95
Β·
verified Β·
1 Parent(s): 73d260a

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +35 -79
app.py CHANGED
@@ -20,12 +20,6 @@ import outlines
20
  from outlines.types import Regex # ε―Όε…₯ζœ€ζ–°ηš„ Regex η±»εž‹ι™εˆΆ
21
  from transformers import AutoTokenizer, AutoModelForCausalLM
22
 
23
- # γ€ζ–°ε’žγ€‘η”Ÿζˆε½“ε‰εˆζ³•θ΅°ζ³•ηš„ζ­£εˆ™θ‘¨θΎΎεΌ
24
- def generate_regex(board: chess.Board) -> str:
25
- legal_moves_san = [board.san(m) for m in board.legal_moves]
26
- pattern = "|".join(re.escape(san) for san in legal_moves_san)
27
- return f" ?({pattern})"
28
-
29
 
30
  # ──────────────────────────────────────────────────────────────────────────────
31
  # Available models
@@ -41,47 +35,28 @@ MODEL_DESCRIPTIONS = {
41
 
42
  device = torch.device("cpu")
43
 
44
- # Lazy cache: {model_id: (tokenizer, model)}
45
- # _model_cache: dict = {}
46
 
47
  @spaces.GPU(duration=30)
48
  def foo(bar):
49
  return bar
50
 
51
- _MODEL_CACHE = {}
52
-
53
 
54
  def load_model(model_key: str):
 
55
  model_id = AVAILABLE_MODELS[model_key]
56
-
57
- # εͺζœ‰η¬¬δΈ€ζ¬‘δ½Ώη”¨θ―₯ζ¨‘εž‹ζ—Άζ‰δΌšεŠ θ½½ζ¨‘εž‹ζƒι‡
58
- if model_id not in _MODEL_CACHE:
59
- print(f"πŸš€ [ι¦–ζ¬‘εŠ θ½½] ζ­£εœ¨εŠ θ½½ζ¨‘εž‹ζƒι‡εˆ°ε†…ε­˜: {model_id} …")
60
- tokenizer = AutoTokenizer.from_pretrained(model_id)
61
- hf_model = AutoModelForCausalLM.from_pretrained(model_id).to(device)
62
-
63
- # δΏε­˜εŒ…θ£…εŽηš„ Outlines ζ¨‘εž‹ε―Ήθ±‘
64
- _MODEL_CACHE[model_id] = outlines.from_transformers(hf_model, tokenizer)
65
- print(f"βœ“ {model_id} ε·²ζˆεŠŸηΌ“ε­˜οΌŒεŽη»­θ΅°ζ£‹ε°†η›΄ζŽ₯ε€η”¨ε†…ε­˜οΌ")
66
-
67
- return _MODEL_CACHE[model_id]
68
- # _model_cache: dict = {}
69
-
70
-
71
- # def load_model(model_key: str):
72
- # """Load (or retrieve from cache) tokenizer + model for the given key."""
73
- # model_id = AVAILABLE_MODELS[model_key]
74
- # if model_id not in _model_cache:
75
- # print(f"Loading {model_id} …")
76
- # tokenizer = GPT2Tokenizer.from_pretrained(model_id)
77
- # tokenizer.pad_token = tokenizer.eos_token
78
- # model = GPT2LMHeadModel.from_pretrained(model_id)
79
- # model.to(device)
80
- # model.eval()
81
- # model.config.use_cache = True
82
- # _model_cache[model_id] = (tokenizer, model)
83
- # print(f"βœ“ {model_id} ready on {device}")
84
- # return _model_cache[model_id]
85
 
86
  # ──────────────────────────────────────────────────────────────────────────────
87
  # Chess / model logic
@@ -130,48 +105,29 @@ def extract_move(text: str, board: chess.Board):
130
 
131
  @torch.no_grad()
132
  def get_model_move(board: chess.Board, model_key: str):
133
- # η›΄ζŽ₯δ»Žε…¨ε±€ηΌ“ε­˜δΈ­θŽ·ε–ε·²εŠ θ½½ε₯½ηš„ζ¨‘εž‹
134
- model = load_model(model_key)
135
-
136
- prompt = board_to_prompt(board)
137
  print("Current prompt: "+prompt)
138
 
139
- regex_pattern = generate_regex(board)
140
-
141
- # η›΄ζŽ₯θΏ›θ‘ŒηΊ¦ζŸη”ŸζˆοΌŒδΈδΌšθ§¦ε‘ζ¨‘εž‹ι‡θ½½
142
- move_san = model(prompt, Regex(regex_pattern), do_sample=True, temperature=0.3, max_new_tokens=10)
143
- move_san = str(move_san).strip()
144
-
145
- move = board.parse_san(move_san)
146
- return move, True
147
-
148
-
149
- # @spaces.GPU(duration=10)
150
- # @torch.no_grad()
151
- # def get_model_move(board: chess.Board, model_key: str):
152
- # tokenizer, model = load_model(model_key)
153
- # prompt = board_to_prompt(board)
154
- # print("Current prompt: "+prompt)
155
- #
156
- # inputs = tokenizer(prompt, return_tensors="pt").to(device)
157
- # outputs = model.generate(
158
- # inputs.input_ids,
159
- # max_new_tokens=12,
160
- # do_sample=True,
161
- # temperature=0.3,
162
- # top_k=40,
163
- # top_p=0.9,
164
- # repetition_penalty=1.1,
165
- # pad_token_id=tokenizer.eos_token_id,
166
- # eos_token_id=tokenizer.eos_token_id,
167
- # )
168
- # new_tokens = outputs[0][inputs.input_ids.shape[1]:]
169
- # generated = tokenizer.decode(new_tokens, skip_special_tokens=True)
170
- # move = extract_move(generated, board)
171
- # if move:
172
- # return move, True
173
- # print(f"Wrong move: {generated}")
174
- # return random.choice(list(board.legal_moves)), False
175
 
176
 
177
 
 
20
  from outlines.types import Regex # ε―Όε…₯ζœ€ζ–°ηš„ Regex η±»εž‹ι™εˆΆ
21
  from transformers import AutoTokenizer, AutoModelForCausalLM
22
 
 
 
 
 
 
 
23
 
24
  # ──────────────────────────────────────────────────────────────────────────────
25
  # Available models
 
35
 
36
  device = torch.device("cpu")
37
 
38
+ Lazy cache: {model_id: (tokenizer, model)}
39
+ _model_cache: dict = {}
40
 
41
  @spaces.GPU(duration=30)
42
  def foo(bar):
43
  return bar
44
 
 
 
45
 
46
  def load_model(model_key: str):
47
+ """Load (or retrieve from cache) tokenizer + model for the given key."""
48
  model_id = AVAILABLE_MODELS[model_key]
49
+ if model_id not in _model_cache:
50
+ print(f"Loading {model_id} …")
51
+ tokenizer = GPT2Tokenizer.from_pretrained(model_id)
52
+ tokenizer.pad_token = tokenizer.eos_token
53
+ model = GPT2LMHeadModel.from_pretrained(model_id)
54
+ model.to(device)
55
+ model.eval()
56
+ model.config.use_cache = True
57
+ _model_cache[model_id] = (tokenizer, model)
58
+ print(f"βœ“ {model_id} ready on {device}")
59
+ return _model_cache[model_id]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
60
 
61
  # ──────────────────────────────────────────────────────────────────────────────
62
  # Chess / model logic
 
105
 
106
  @torch.no_grad()
107
  def get_model_move(board: chess.Board, model_key: str):
108
+ tokenizer, model = load_model(model_key)
109
+ prompt = board_to_prompt(board)
 
 
110
  print("Current prompt: "+prompt)
111
 
112
+ inputs = tokenizer(prompt, return_tensors="pt").to(device)
113
+ outputs = model.generate(
114
+ inputs.input_ids,
115
+ max_new_tokens=12,
116
+ do_sample=True,
117
+ temperature=0.3,
118
+ top_k=40,
119
+ top_p=0.9,
120
+ repetition_penalty=1.1,
121
+ pad_token_id=tokenizer.eos_token_id,
122
+ eos_token_id=tokenizer.eos_token_id,
123
+ )
124
+ new_tokens = outputs[0][inputs.input_ids.shape[1]:]
125
+ generated = tokenizer.decode(new_tokens, skip_special_tokens=True)
126
+ move = extract_move(generated, board)
127
+ if move:
128
+ return move, True
129
+ print(f"Wrong move: {generated}")
130
+ return random.choice(list(board.legal_moves)), False
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
131
 
132
 
133