File size: 13,667 Bytes
216e5d3
565c3f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
afb6483
0196d48
565c3f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
import os
# Expandable segments guards against transient allocator fragmentation on the
# ~23 GB checkpoint + activations.
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")

import sys
from pathlib import Path

# Vendored custom architecture code (qwenvl package) lives under ./code
CODE_DIR = Path(__file__).parent / "code"
sys.path.insert(0, str(CODE_DIR))

import spaces  # noqa: E402  MUST come before torch / CUDA-touching imports
import torch  # noqa: E402
import gradio as gr  # noqa: E402
from huggingface_hub import snapshot_download  # noqa: E402

MODEL_ID = "sais-org/Polaris_Pro"

# ---------------------------------------------------------------------------
# Download weights at module scope (pure I/O, no CUDA). We deliberately do NOT
# import the model code here: the mol modality pulls in `torch_geometric`,
# whose import initialises a real CUDA context in the process. On ZeroGPU that
# poisons the forked GPU worker ("No CUDA GPUs are available" in worker_init).
# So the heavy `import inference` + model construction + move-to-cuda all run
# lazily INSIDE the @spaces.GPU worker (where a real GPU is attached), and the
# built model is cached in a module global for subsequent (warm) calls.
# ---------------------------------------------------------------------------
print("Downloading model weights ...")
MODEL_DIR = snapshot_download(
    MODEL_ID,
    token=os.environ.get("HF_TOKEN"),
)
print(f"Model downloaded to {MODEL_DIR}")

INFER = None


def _ensure_model():
    """Build BioQwen3VLInference on the GPU exactly once, inside the worker.

    Importing here (not at module scope) keeps torch_geometric's CUDA init out
    of the main process. Cached globally so warm workers skip the rebuild.
    """
    global INFER
    if INFER is not None:
        return INFER
    from inference import BioQwen3VLInference
    print("Instantiating BioQwen3VLInference on CUDA (~23 GB of weights) ...")
    INFER = BioQwen3VLInference(
        model_path=MODEL_DIR,
        device="cuda",
        dtype=torch.bfloat16,
        attn_impl="sdpa",          # torch-native; correct on ZeroGPU Blackwell
        fail_on_legacy_mol_decoder=False,
    )
    print("Model ready (on CUDA).")
    return INFER


# ---------------------------------------------------------------------------
# Per-task presets: system prompt + prompt template + which modality field the
# sequence input maps to. Prompts / sequences mirror run_examples.sh, which the
# authors ship to reproduce the benchmark numbers.
# ---------------------------------------------------------------------------
TASKS = {
    "RNA · ncRNA family classification": {
        "field": "rna",
        "system": "You are a non-coding RNA family classifier. Output only the family name, no other text.",
        "prompt": "<rna>\nWhich family does this non-coding RNA sequence belong to?",
        "task": None,
        "placeholder": "RNA / cDNA nucleotide sequence (A/C/G/U or A/C/G/T)",
    },
    "RNA · translation efficiency (regression)": {
        "field": "rna",
        "system": None,
        "prompt": "<rna>\nWhat is the expected translation efficiency associated with the sequence?",
        "task": None,
        "placeholder": "RNA nucleotide sequence",
    },
    "DNA · promoter detection (Yes/No)": {
        "field": "dna",
        "system": "You are a DNA sequence analysis expert. Read the DNA sequence(s) and the question carefully. Respond with a single token: exactly 'Yes' or 'No'. Do not add any explanation, punctuation, reasoning, or additional text.",
        "prompt": "<dna>\nIs this 300 bp DNA sequence a promoter region (all promoters, TATA and non-TATA combined)? Answer Yes or No.",
        "task": None,
        "placeholder": "DNA nucleotide sequence (A/C/G/T)",
    },
    "DNA · enhancer activity (regression)": {
        "field": "dna",
        "system": "You are a DNA sequence analysis expert. Read the DNA sequence and the question carefully. Respond with a single floating-point number only. Do not add units, explanations, reasoning, or any additional text.",
        "prompt": "<dna>\nPredict the quantile-normalized developmental enhancer (Dev) log2 enrichment activity score of this DNA sequence. Answer with a float number.",
        "task": None,
        "placeholder": "DNA nucleotide sequence (A/C/G/T)",
    },
    "Protein · solubility (0/1)": {
        "field": "protein",
        "system": "You are a protein solubility predictor. This is a binary classification task. Output only one digit: 1 for soluble, 0 for insoluble. Do not output any other text.",
        "prompt": "<protein>\nSolubility prediction involves forecasting if a protein can dissolve. What is the solubility status of this protein? Output only one digit: 1 for soluble, 0 for insoluble.",
        "task": None,
        "placeholder": "Protein amino-acid sequence",
    },
    "Protein · stability (regression)": {
        "field": "protein",
        "system": "You are a protein stability predictor. Output only the stability score as a number, no other text.",
        "prompt": "<protein>\nHow is the stability of this protein sequence calculated?",
        "task": None,
        "placeholder": "Protein amino-acid sequence",
    },
    "Protein · Enzyme Commission number": {
        "field": "protein",
        "system": "You are a protein function predictor. Output only the EC number(s), comma-separated, no other text.",
        "prompt": "<protein>\nPredict the Enzyme Commission (EC) number(s) of this protein. Output only the EC numbers, comma-separated.",
        "task": None,
        "placeholder": "Protein amino-acid sequence",
    },
    "Molecule · Ames mutagenicity (0/1)": {
        "field": "mol",
        "system": "You are a molecular property prediction expert; given a molecule's SMILES string and an ADMET endpoint description, respond with only 0 or 1 to indicate whether the molecule possesses that property.",
        "prompt": "<mol>\nGiven the SMILES representation of a molecule, predict whether it is mutagenic (1) or non-mutagenic (0) based on the Ames test.",
        "task": None,
        "placeholder": "Molecule SMILES string",
    },
    "Molecule · dipole moment (regression)": {
        "field": "mol",
        "system": "You are a molecular property prediction expert. Based on the input molecular representations and instructions, answer with the specific molecular property values.",
        "prompt": "<mol>\nWhat is the dipole moment value of this molecular?",
        "task": None,
        "placeholder": "Molecule SMILES string",
    },
    "Molecule · text → SMILES (generation)": {
        "field": "text",   # no bio input; description goes in the prompt
        "system": "You are a molecule generation expert. Given a natural-language molecular description, generate one molecule as a valid canonical SMILES string. Output only the SMILES string, with no additional text.",
        "prompt": "Generate a molecule that matches the following description:\n{input}\nOutput only the canonical SMILES string.",
        "task": "mol_generation",
        "placeholder": "Natural-language description of the molecule to generate",
    },
    "Scientific text QA (no sequence)": {
        "field": "text",
        "system": None,
        "prompt": "{input}",
        "task": None,
        "placeholder": "Ask a scientific question (multiple-choice, definitions, reasoning, ...)",
    },
}

TASK_NAMES = list(TASKS.keys())


def _duration(task_name, seq_input, max_new_tokens=64, *args, **kwargs):
    # A cold worker pays the one-time build+load of the ~23 GB model inside the
    # fork (from_pretrained disk->GPU, see _ensure_model). Budget for the worst
    # case; warm workers finish in a fraction of this.
    n = int(max_new_tokens or 64)
    if INFER is None:                 # cold worker: model not yet built
        return min(300, 180 + int(n * 0.8))
    return min(160, 40 + int(n * 0.8))


@spaces.GPU(duration=_duration)
def run_inference(task_name: str, seq_input: str, max_new_tokens: int = 64) -> str:
    """Run Polaris-Pro on one scientific input and return the model's text answer.

    Args:
        task_name: Which scientific task / output format to use (see the dropdown).
        seq_input: The biological sequence (RNA/DNA/protein SMILES) or, for
            text tasks, the natural-language question / molecule description.
        max_new_tokens: Maximum number of tokens to generate.

    Returns:
        The model's text response (a label, score, SMILES string, or answer).
    """
    if task_name not in TASKS:
        return "Unknown task."
    infer = _ensure_model()
    spec = TASKS[task_name]
    seq_input = (seq_input or "").strip()
    if not seq_input:
        return "Please provide an input sequence / question."

    field = spec["field"]
    system = spec["system"]
    prompt_tmpl = spec["prompt"]
    task = spec["task"]

    kwargs = dict(
        max_new_tokens=int(max_new_tokens),
        do_sample=False,          # greedy — matches run_examples.sh (--greedy)
        system=system,
        task=task,
    )

    if field == "text":
        prompt = prompt_tmpl.format(input=seq_input) if "{input}" in prompt_tmpl else prompt_tmpl
        kwargs["prompt"] = prompt
    else:
        # Bio-sequence task: the sequence goes to its own encoder field, and
        # the prompt already carries the matching <rna>/<dna>/... placeholder.
        seq = seq_input.upper() if field in ("rna", "dna", "protein") else seq_input
        kwargs["prompt"] = prompt_tmpl
        kwargs[field] = [seq]

    out = infer.generate_from_prompt(**kwargs)
    return (out or "").strip() or "(empty response)"


def on_task_change(task_name):
    spec = TASKS.get(task_name, {})
    return gr.update(placeholder=spec.get("placeholder", ""))


# ---------------------------------------------------------------------------
# Example inputs (task_name, sequence/question, max_new_tokens) — taken from
# the authors' run_examples.sh.
# ---------------------------------------------------------------------------
EXAMPLES = [
    ["RNA · ncRNA family classification",
     "GGATGCGATCATGTCTGCACTAACACACCGGATCCCATCAGAACTCCGAAGTTAAGCGTGCTTGGGCGGGAGTAGTACTAGGATGGGCGACCCCTTAGGAAGTACTCGTGTTGCATCCC",
     64],
    ["DNA · promoter detection (Yes/No)",
     "GCAATAAAAGGCTTAGCCACATAGTGCATGCATGTACACAGCATGTACAC",
     16],
    ["Protein · solubility (0/1)",
     "MLSVRIAAAVARALPRRAGLVSKNALGSSFIAARNFHASNTHLQKTGTAEMSSILEERILGADTSVDLEETGRVLSIGDGIARVHGLRNVQAEEMVEFSSGLKGMSLNLEP",
     16],
    ["Protein · Enzyme Commission number",
     "MHHHHHHSSGVDLGTENLYFQSNAMDFPQQLEACVKQANQALSRFIAPLPFQNTPVVETMQYGALLGGKRLRPFLVYATGHMFGVSTNTLDAPAAAVELIHAYSLIHDDLPAMDDDDLRRGLPTCHVKFGEANAILAGDALQTLAFSILSDADLADYIIQRNK",
     32],
    ["Molecule · Ames mutagenicity (0/1)",
     "CC(=O)Nc1ccc2c(=O)c(=O)c3cccc4ccc1c2c43",
     16],
    ["Molecule · text → SMILES (generation)",
     "The molecule is a long-chain fatty acid that is henicosane in which one of the methyl groups has been oxidised to give the corresponding carboxylic acid. It is a straight-chain saturated fatty acid and a long-chain fatty acid.",
     128],
    ["Scientific text QA (no sequence)",
     "The following is a multiple choice question about biology. Think step by step and then finish your answer with \"the answer is (X)\".\nQuestion:\nWhich molecule carries amino acids to the ribosome during translation?\nOptions:\nA. mRNA\nB. tRNA\nC. rRNA\nD. snRNA\nAnswer:",
     256],
]

CSS = """
#col-container { max-width: 1080px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""

with gr.Blocks() as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(
            """
            # 🔬 Polaris-Pro — Unified Scientific Multimodal Model

            [`sais-org/Polaris_Pro`](https://huggingface.co/sais-org/Polaris_Pro) is an **8B** foundation
            model that reasons over proteins, RNA, DNA, and small molecules through a single
            natural-language interface — no per-task fine-tuning.

            Pick a task, paste a sequence (or a question), and run. Each task uses the authors'
            official system prompt so the output format matches their benchmarks.

            *Weather forecasting and medical-image segmentation are part of the model but need
            gridded netCDF I/O / gated SAM-3 weights, so they are out of scope for this demo.*
            """
        )

        with gr.Row():
            task = gr.Dropdown(
                choices=TASK_NAMES, value=TASK_NAMES[0], label="Task", scale=2,
            )
            run = gr.Button("Run", variant="primary", scale=1)

        seq = gr.Textbox(
            label="Input",
            lines=4,
            placeholder=TASKS[TASK_NAMES[0]]["placeholder"],
        )
        output = gr.Textbox(label="Model response", lines=4)

        with gr.Accordion("Advanced settings", open=False):
            max_new = gr.Slider(
                label="Max new tokens", minimum=1, maximum=512, value=64, step=1,
            )

        gr.Examples(
            examples=EXAMPLES,
            inputs=[task, seq, max_new],
            outputs=output,
            fn=run_inference,
            cache_examples=False,
            run_on_click=True,
        )

    task.change(on_task_change, inputs=task, outputs=seq)
    run.click(run_inference, inputs=[task, seq, max_new], outputs=output, api_name="generate")

if __name__ == "__main__":
    demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)