unruffle commited on
Commit
efab3ba
·
verified ·
1 Parent(s): 85bdae4

Upload 6 files

Browse files
Files changed (6) hide show
  1. Dockerfile +52 -0
  2. README.md +91 -5
  3. app.py +268 -0
  4. gitattributes +35 -0
  5. hf_loader.py +40 -0
  6. requirements.txt +13 -0
Dockerfile ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.10-slim
2
+
3
+ # 1. 环境变量
4
+ ENV DEBIAN_FRONTEND=noninteractive \
5
+ PYTHONDONTWRITEBYTECODE=1 \
6
+ PYTHONUNBUFFERED=1 \
7
+ PIP_NO_CACHE_DIR=1 \
8
+ HF_HOME=/app/cache/huggingface \
9
+ PYTHONIOENCODING=UTF-8 \
10
+ PORT=7860 \
11
+ MODEL_ID="knowledgator/SMILES2IUPAC-canonical-base"
12
+
13
+ # 2. 系统依赖
14
+ RUN apt-get update && apt-get install -y --no-install-recommends \
15
+ curl ca-certificates \
16
+ libstdc++6 libgomp1 \
17
+ libxrender1 libxext6 \
18
+ tini \
19
+ && rm -rf /var/lib/apt/lists/*
20
+
21
+ WORKDIR /app
22
+
23
+ # 3. 用户权限
24
+ RUN useradd -m -u 1000 user
25
+
26
+ # 4. Python 依赖
27
+ COPY requirements.txt /app/requirements.txt
28
+ RUN pip install --upgrade pip && \
29
+ pip install -r requirements.txt
30
+
31
+ # 5. 预下载模型
32
+ # 必须通过 NamesConverter 加载(模型使用自定义小词表,不兼容 AutoModelForSeq2SeqLM)
33
+ # 同时打印内部属性名,方便调试确认 hf_loader.py 中的属性探测是否正确
34
+ RUN python -c "\
35
+ from chemicalconverters import NamesConverter; \
36
+ import os; \
37
+ model_id = os.getenv('MODEL_ID', 'knowledgator/SMILES2IUPAC-canonical-base'); \
38
+ print(f'Pre-loading: {model_id}'); \
39
+ c = NamesConverter(model_name=model_id); \
40
+ attrs = [a for a in dir(c) if not a.startswith('__')]; \
41
+ print(f'NamesConverter attrs: {attrs}'); \
42
+ print('Pre-load complete.')"
43
+
44
+ # 6. 复制文件
45
+ COPY --chown=user:user . /app
46
+ RUN mkdir -p /app/cache/huggingface && chown -R user:user /app/cache
47
+
48
+ # 7. 启动
49
+ USER user
50
+ EXPOSE $PORT
51
+ ENTRYPOINT ["/usr/bin/tini", "--"]
52
+ CMD ["sh", "-c", "uvicorn app:app --host 0.0.0.0 --port ${PORT:-7860}"]
README.md CHANGED
@@ -1,10 +1,96 @@
1
  ---
2
- title: MolScribe
3
- emoji: 👁
4
- colorFrom: purple
5
- colorTo: red
6
  sdk: docker
7
  pinned: false
8
  ---
 
9
 
10
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: SMILES → IUPAC (Docker)
3
+ emoji: 🧪
4
+ colorFrom: indigo
5
+ colorTo: blue
6
  sdk: docker
7
  pinned: false
8
  ---
9
+ # SMILES→IUPAC(Attention-Optimized V6.6)
10
 
11
+ 基于`knowledgator/SMILES2IUPAC-canonical-base`模型的SMILES转IUPAC高精度转换服务。本项目集成了FastAPI与Gradio,引入了基于注意力的跨度预优化、动态束搜索与拓扑幻觉拦截机制,完美兼顾了简单芳香环的高效解码与复杂手性骨架的深度探索。
12
+
13
+ ## 核心特性
14
+
15
+ * **动态束搜索(Dynamic Beam Search)**:系统会通过RDKit实时计算分子的重原子数量(Heavy Atoms)。小分子(如苯环,≤10)自动采用贪婪解码(`beams=1`)彻底根治同位素幻觉;中型分子(≤20)使用`beams=4`平衡速度;大分子(如福莫特罗,>20)自动开启深度搜索(`beams=8`)防止长序列推理断连。
16
+ * **TTA短路验证机制(Short-circuit Evaluation)**:在处理复杂大分子时,系统会生成包含50个合规变体的候选池,并依次送入大语言模型进行推理。一旦某个变体生成的IUPAC名称通过了拓扑安全校验,即刻短路跳出并返回结果,极大节省了显存算力与API响应时间。
17
+ * **纯随机遍历预优化**:彻底废除了破坏底层芳香属性的强制Kekulize操作,恢复极具多样性的纯随机遍历(`doRandom=True`)算法,并配合生成大写凯库勒式(`kekuleSmiles=True`),在最大化变体结构多样性的同时,完美契合大模型底层的无芳香标志预训练分布。
18
+ * **RDKit拓扑拦截过滤**:基于真实分子拓扑结构,自动拦截并剔除AI生成的含有不存在环系(如`cyclohept`、`cyclooct`、`cyclonon`)的虚假名称。
19
+ * **双重交互模式**:提供标准化RESTful API供程序调用,同时挂载Gradio WebUI方便直观调试与查看变体推理细节。
20
+
21
+ ## 环境变量配置
22
+
23
+ * `DISABLE_CANONICALIZE`:设为`1`或`true`可全局关闭服务端强制规范化(不推荐,关闭后若输入非标SMILES可能诱发幻觉)。
24
+ * `TTA_SAMPLES`:TTA候选池的评估上限,开启`use_tta`时生效,默认最少评估`5`个优质变体。
25
+ * `MODEL_ID`:底层模型路径,默认为`knowledgator/SMILES2IUPAC-canonical-base`。
26
+
27
+ ## REST API调用指南
28
+
29
+ ### 1.服务健康检查
30
+ **GET** `/healthz`
31
+ 返回结果:
32
+ ```json
33
+ {
34
+ "ok": true
35
+ }
36
+
37
+ ```
38
+
39
+ ### 2.单条SMILES转换
40
+
41
+ **POST** `/api/smiles2iupac`
42
+ **请求体(JSON):**
43
+
44
+ ```json
45
+ {
46
+ "smiles": "COc1ccc(CC(C)NCC(O)c2ccc(O)c(NC=O)c2)cc1",
47
+ "canonicalize": true,
48
+ "style": "BASE",
49
+ "use_tta": true
50
+ }
51
+
52
+ ```
53
+
54
+ **字段说明:**
55
+
56
+ * `smiles`:需要转换的SMILES字符串。
57
+ * `canonicalize`:是否允许服务端进行标准化处理(强烈建议为`true`)。
58
+ * `style`:命名风格,支持`BASE`(默认推荐,兼顾俗名与系统名)、`SYST`(纯系统命名)、`TRAD`(传统命名)。
59
+ * `use_tta`:是否开启跨度预优化与短路验证机制。
60
+
61
+ **响应体(JSON):**
62
+
63
+ ```json
64
+ {
65
+ "success": true,
66
+ "input": "COc1ccc(CC(C)NCC(O)c2ccc(O)c(NC=O)c2)cc1",
67
+ "style": "BASE",
68
+ "tta_used": true,
69
+ "candidates_count": 5,
70
+ "iupac": "N-[2-hydroxy-5-[1-hydroxy-2-[[1-(4-methoxyphenyl)propan-2-yl]amino]ethyl]phenyl]formamide",
71
+ "voting_details": {
72
+ "N-[2-hydroxy-5-[1-hydroxy-2-[[1-(4-methoxyphenyl)propan-2-yl]amino]ethyl]phenyl]formamide": 1
73
+ }
74
+ }
75
+
76
+ ```
77
+
78
+ ### 3.批量SMILES转换
79
+
80
+ **POST** `/api/smiles2iupac/batch`
81
+ **请求体(JSON):**
82
+
83
+ ```json
84
+ {
85
+ "inputs": [
86
+ {
87
+ "smiles": "c1ccccc1",
88
+ "style": "BASE"
89
+ },
90
+ {
91
+ "smiles": "CC(=O)Oc1ccccc1C(=O)O",
92
+ "style": "SYST",
93
+ "use_tta": true
94
+ }
95
+ ]
96
+ }
app.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # -*- coding: utf-8 -*-
2
+ import os
3
+ import re
4
+ from collections import Counter
5
+ import gradio as gr
6
+ from fastapi import FastAPI
7
+ from fastapi.middleware.cors import CORSMiddleware
8
+ from pydantic import BaseModel
9
+ from typing import List, Optional, Union
10
+
11
+ from rdkit import Chem
12
+ from rdkit.Chem import MolToSmiles # 彻底移除了破坏性的Kekulize
13
+
14
+ from hf_loader import smiles2iupac
15
+
16
+ app = FastAPI(
17
+ title="SMILES→IUPAC(Attention-Optimized)",
18
+ version="6.6.0",
19
+ )
20
+
21
+ app.add_middleware(
22
+ CORSMiddleware,
23
+ allow_origins=["*"],
24
+ allow_methods=["*"],
25
+ allow_headers=["*"],
26
+ )
27
+
28
+ def calculate_max_ring_span(smiles: str) -> int:
29
+ def repl(m):
30
+ return '_' * len(m.group())
31
+
32
+ clean_smiles = re.sub(r'\[.*?\]', repl, smiles)
33
+
34
+ active_rings = {}
35
+ max_span = 0
36
+
37
+ for match in re.finditer(r'%\d{2}|\d', clean_smiles):
38
+ ring_id = match.group()
39
+ pos = match.start()
40
+
41
+ if ring_id in active_rings:
42
+ span = pos - active_rings[ring_id]
43
+ if span > max_span:
44
+ max_span = span
45
+ del active_rings[ring_id]
46
+ else:
47
+ active_rings[ring_id] = pos
48
+
49
+ return max_span
50
+
51
+ def generate_optimized_smiles(s: str, pool_size: int = 50, top_k: int = 5) -> List[str]:
52
+ mol = Chem.MolFromSmiles(s)
53
+ if mol is None:
54
+ raise ValueError(f"非法SMILES,RDKit无法解析:{s!r}")
55
+
56
+ # 获取标准规范凯库勒式作为保底(RDKit内部会自动处理,不破坏mol对象)
57
+ canonical_kekule = MolToSmiles(mol, canonical=True, kekuleSmiles=True)
58
+ num_atoms = mol.GetNumAtoms()
59
+
60
+ if num_atoms < 15:
61
+ return [canonical_kekule]
62
+
63
+ tta_set = {canonical_kekule}
64
+
65
+ # 核心修复:恢复极其成功的随机算法doRandom=True
66
+ max_attempts = pool_size * 5
67
+ attempts = 0
68
+ while len(tta_set) < pool_size and attempts < max_attempts:
69
+ try:
70
+ variant = MolToSmiles(mol, canonical=False, doRandom=True, kekuleSmiles=True)
71
+ tta_set.add(variant)
72
+ except Exception:
73
+ pass
74
+ attempts += 1
75
+
76
+ scored_smiles = []
77
+ for sm in tta_set:
78
+ span_score = calculate_max_ring_span(sm)
79
+ scored_smiles.append((span_score, len(sm), sm))
80
+
81
+ scored_smiles.sort()
82
+ return [item[2] for item in scored_smiles[:top_k]]
83
+
84
+
85
+ DISABLE_CANONICALIZE = os.getenv("DISABLE_CANONICALIZE", "").lower() in ("1", "true", "yes")
86
+ TTA_SAMPLES = int(os.getenv("TTA_SAMPLES", "5"))
87
+
88
+ class SMILESItem(BaseModel):
89
+ smiles: str
90
+ canonicalize: Optional[bool] = True
91
+ style: Optional[str] = "BASE"
92
+ use_tta: Optional[bool] = True
93
+ beams: Optional[Union[int, str]] = "auto"
94
+
95
+ class BatchRequest(BaseModel):
96
+ inputs: List[SMILESItem]
97
+
98
+ def process_single_smiles(s: str, do_canon: bool, style: str, use_tta: bool, beams: Union[int, str] = "auto"):
99
+ mol = Chem.MolFromSmiles(s)
100
+ if mol is None:
101
+ return "", [s], {}, 1
102
+
103
+ if str(beams).lower() == "auto":
104
+ heavy_atoms = mol.GetNumHeavyAtoms()
105
+ if heavy_atoms <= 10:
106
+ dynamic_beams = 1
107
+ elif heavy_atoms <= 20:
108
+ dynamic_beams = 4
109
+ else:
110
+ dynamic_beams = 8
111
+ else:
112
+ try:
113
+ dynamic_beams = int(beams)
114
+ if dynamic_beams < 1:
115
+ dynamic_beams = 1
116
+ except ValueError:
117
+ dynamic_beams = 4
118
+
119
+ ring_info = mol.GetRingInfo().AtomRings()
120
+ ring_sizes = set(len(r) for r in ring_info)
121
+
122
+ hallucination_blacklist = []
123
+ if 7 not in ring_sizes: hallucination_blacklist.append("cyclohept")
124
+ if 8 not in ring_sizes: hallucination_blacklist.append("cyclooct")
125
+ if 9 not in ring_sizes: hallucination_blacklist.append("cyclonon")
126
+
127
+ if not do_canon:
128
+ try:
129
+ # 同样移除对mol对象的破坏,直接输出大写
130
+ s_kekule = Chem.MolToSmiles(mol, canonical=False, kekuleSmiles=True)
131
+ except Exception:
132
+ s_kekule = s
133
+
134
+ name = smiles2iupac(s_kekule, style=style, num_beams=dynamic_beams)
135
+ return name, [s_kekule], {name: 1}, dynamic_beams
136
+
137
+ sample_count = TTA_SAMPLES if use_tta else 1
138
+ if use_tta and sample_count < 5:
139
+ sample_count = 5
140
+
141
+ smiles_list = generate_optimized_smiles(s, pool_size=50, top_k=sample_count)
142
+
143
+ evaluated_smiles = []
144
+ best_name = ""
145
+
146
+ for current_smiles in smiles_list:
147
+ evaluated_smiles.append(current_smiles)
148
+ name = smiles2iupac(current_smiles, style=style, num_beams=dynamic_beams)
149
+
150
+ if not name:
151
+ continue
152
+
153
+ name_lower = name.lower()
154
+ if any(bad_word in name_lower for bad_word in hallucination_blacklist):
155
+ continue
156
+
157
+ best_name = name
158
+ break
159
+
160
+ if not best_name:
161
+ return "未能生成符合拓结构的名称(全被过滤器拦截)", evaluated_smiles, {}, dynamic_beams
162
+
163
+ return best_name, evaluated_smiles, {best_name: 1}, dynamic_beams
164
+
165
+ @app.get("/healthz")
166
+ def healthz():
167
+ return {"ok": True}
168
+
169
+ @app.post("/api/smiles2iupac")
170
+ def api_smiles2iupac(req: SMILESItem):
171
+ try:
172
+ s = (req.smiles or "").strip()
173
+ if not s:
174
+ return {"success": False, "error": "输入为空"}
175
+
176
+ do_canon = (req.canonicalize if req.canonicalize is not None else True) and (not DISABLE_CANONICALIZE)
177
+ valid_styles = ["BASE", "SYST", "TRAD"]
178
+ style = req.style.upper() if req.style and req.style.upper() in valid_styles else "BASE"
179
+ use_tta = req.use_tta if req.use_tta is not None else True
180
+ beams = req.beams if req.beams is not None else "auto"
181
+
182
+ best_name, smiles_list, counts, _ = process_single_smiles(s, do_canon, style, use_tta, beams)
183
+
184
+ return {
185
+ "success": True,
186
+ "input": s,
187
+ "style": style,
188
+ "tta_used": use_tta,
189
+ "candidates_count": len(smiles_list),
190
+ "iupac": best_name,
191
+ "voting_details": counts
192
+ }
193
+ except Exception as e:
194
+ return {"success": False, "error": str(e)}
195
+
196
+ @app.post("/api/smiles2iupac/batch")
197
+ def api_smiles2iupac_batch(req: BatchRequest):
198
+ out = []
199
+ for item in req.inputs:
200
+ try:
201
+ s = (item.smiles or "").strip()
202
+ if not s:
203
+ out.append({"success": False, "error": "输入为空"})
204
+ continue
205
+
206
+ do_canon = (item.canonicalize if item.canonicalize is not None else True) and (not DISABLE_CANONICALIZE)
207
+ valid_styles = ["BASE", "SYST", "TRAD"]
208
+ style = item.style.upper() if item.style and item.style.upper() in valid_styles else "BASE"
209
+ use_tta = item.use_tta if item.use_tta is not None else True
210
+ beams = item.beams if item.beams is not None else "auto"
211
+
212
+ best_name, smiles_list, counts, _ = process_single_smiles(s, do_canon, style, use_tta, beams)
213
+
214
+ out.append({
215
+ "success": True,
216
+ "input": s,
217
+ "style": style,
218
+ "iupac": best_name
219
+ })
220
+ except Exception as e:
221
+ out.append({"success": False, "input": item.smiles, "error": str(e)})
222
+ return out
223
+
224
+ def gradio_fn(s: str, style: str, canonicalize: bool, use_tta: bool, beams: str):
225
+ if not (s or "").strip():
226
+ return "", "输入为空"
227
+ try:
228
+ do_canon = canonicalize and (not DISABLE_CANONICALIZE)
229
+ best_name, smiles_list, counts, dynamic_beams = process_single_smiles(s, do_canon, style, use_tta, beams)
230
+
231
+ if not do_canon and not use_tta:
232
+ debug_info = (
233
+ f"原始输入:{s}\n"
234
+ f"[直通模式]:动态BeamSearch宽度:{dynamic_beams}\n\n"
235
+ f"模型直接输出:{best_name}"
236
+ )
237
+ else:
238
+ variants_text = "\n".join([f" - {x}" for x in smiles_list])
239
+ attempts_count = len(smiles_list)
240
+
241
+ debug_info = (
242
+ f"原始输入:{s}\n"
243
+ f"动态BeamSearch宽度:{dynamic_beams}\n\n"
244
+ f"执行短路验证(大模型实际推理了{attempts_count}个高质量变体):\n{variants_text}\n\n"
245
+ f"最终采纳结果(验证通过):\n - {best_name}\n"
246
+ )
247
+ return best_name, debug_info
248
+ except Exception as e:
249
+ return "", f"Error:{e}"
250
+
251
+ demo = gr.Interface(
252
+ fn=gradio_fn,
253
+ inputs=[
254
+ gr.Textbox(label="输入SMILES", placeholder="支持任意合法SMILES写法"),
255
+ gr.Radio(["BASE", "SYST", "TRAD"], label="命名风格", value="BASE"),
256
+ gr.Checkbox(label="自动Kekulé规范化(强烈建议开启)", value=True),
257
+ gr.Checkbox(label="开启跨度预优化与拓扑过滤(解决复杂环系幻觉)", value=True),
258
+ gr.Dropdown(["auto", "1", "4", "8", "16"], label="BeamSearch宽度(选auto为智能调整)", value="auto", allow_custom_value=True),
259
+ ],
260
+ outputs=[
261
+ gr.Textbox(label="最优IUPAC名称", interactive=True),
262
+ gr.Textbox(label="调试信息", lines=12),
263
+ ],
264
+ title="SMILES→IUPAC(Attention-Optimized V6.6)",
265
+ description="恢复基于纯随机遍历的变体生成算法以最大化结构多样性;配合动态束搜索与短路拦截,实现精度与速度的双赢。",
266
+ )
267
+
268
+ app = gr.mount_gradio_app(app, demo, path="/")
gitattributes ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz 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
hf_loader.py ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # -*- coding: utf-8 -*-
2
+ import os
3
+ from functools import lru_cache
4
+ from typing import Union, List
5
+ from chemicalconverters import NamesConverter
6
+
7
+ MODEL_ID = os.getenv("MODEL_ID", "knowledgator/SMILES2IUPAC-canonical-base")
8
+
9
+ @lru_cache(maxsize=1)
10
+ def _load() -> NamesConverter:
11
+ print(f"[hf_loader]Loading:{MODEL_ID}...")
12
+ converter = NamesConverter(
13
+ model_name=MODEL_ID,
14
+ smiles_max_len=512,
15
+ iupac_max_len=512
16
+ )
17
+ print("[hf_loader]Loaded.")
18
+ return converter
19
+
20
+ def smiles2iupac(smiles_kekule: Union[str, List[str]], style: str = "BASE", num_beams: int = 4) -> Union[str, List[str]]:
21
+ """
22
+ 支持单条或批量SMILES转换为IUPAC名称,支持动态束搜索宽度
23
+ """
24
+ valid_styles = {"BASE", "SYST", "TRAD"}
25
+ style_upper = style.upper() if style and style.upper() in valid_styles else "BASE"
26
+
27
+ converter = _load()
28
+ is_list = isinstance(smiles_kekule, list)
29
+ smiles_list = smiles_kekule if is_list else [smiles_kekule]
30
+
31
+ results = []
32
+ for s in smiles_list:
33
+ input_text = f"<{style_upper}>{s}"
34
+ res = converter.smiles_to_iupac(input_text, num_beams=num_beams)
35
+ if isinstance(res, list):
36
+ results.append(res[0] if res else "")
37
+ else:
38
+ results.append(str(res))
39
+
40
+ return results if is_list else results[0]
requirements.txt ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ fastapi
2
+ uvicorn
3
+ gradio>=4.0
4
+ huggingface_hub
5
+ # 官方库 — 必须保留,因为模型使用自定义小词表(encoder:137 / decoder:822 tokens)
6
+ # 只有通过 NamesConverter 才能正确加载这两个自定义 tokenizer
7
+ # 我们在 hf_loader.py 中绕过了其内部有 bug 的预处理,直接操控底层推理
8
+ chemical-converters>=0.1.2
9
+ transformers>=4.35.0
10
+ torch --extra-index-url https://download.pytorch.org/whl/cpu
11
+ sentencepiece
12
+ protobuf
13
+ rdkit-pypi