#!/usr/bin/env python """ Unified inference script for Polaris-Pro. Supported tasks (matching training format): - classification : RNA/DNA/mol input -> text label (0/1) - regression : RNA+DNA/mol input -> score - description : RNA input -> free-form text - generation : RNA input(s) -> RNA output sequence - mol_property_* : mol input -> property prediction Input JSON(L) format is identical to training: { "rna": ["SEQ1", ...], # optional "dna": ["SEQ1", ...], # optional "mol": ["SMILES1", ...], # optional "image": ["path.jpg"], # optional "conversations": [ {"from": "human", "value": "\n...prompt..."}, {"from": "gpt", "value": ""} # leave empty for inference ], "task": "mol_property_classification" # optional, for logging } Usage: # Single sample (inline) python inference.py --model_path /path/to/ckpt \ --rna "AUGCAUGC" \ --prompt "\nWhat family does this RNA belong to?" # Batch from file python inference.py --model_path /path/to/ckpt \ --input_file samples.jsonl \ --output_file results.jsonl # Interactive chat python inference.py --model_path /path/to/ckpt --chat """ import argparse import json import re import sys from pathlib import Path from typing import Dict, List, Optional, Any, Tuple import torch from PIL import Image project_root = Path(__file__).parent sys.path.insert(0, str(project_root)) try: from transformers.generation.logits_process import LogitsProcessor, LogitsProcessorList except ImportError: # transformers compatibility shim from transformers import LogitsProcessor, LogitsProcessorList from qwenvl.models import Qwen3VLForConditionalGeneration from qwenvl.processor_compat import load_auto_processor_compat from qwenvl.data.inference_helpers import ( RNACharTokenizer, RNA_NUM_LATENT_TOKENS, _RNA_CHAR_TO_ID, _RNA_ID_TO_CHAR, extract_epi_cell_line, ea_n_bins_for_sample, rewrite_ea_binned_instruction, ) from qwenvl.registry.token_manager import BIO_SEQ_OUTPUT_PAD, MODALITY_TOKEN_DEFS _RNA_TOKENS = MODALITY_TOKEN_DEFS["rna"] _DNA_TOKENS = MODALITY_TOKEN_DEFS["dna"] _PROTEIN_TOKENS = MODALITY_TOKEN_DEFS["protein"] _MOL_TOKENS = MODALITY_TOKEN_DEFS["mol"] class PresencePenaltyLogitsProcessor(LogitsProcessor): """Apply vLLM/OpenAI-style presence penalty to generated tokens only.""" def __init__(self, prompt_length: int, penalty: float): self.prompt_length = int(prompt_length) self.penalty = float(penalty) def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor: if self.penalty <= 0 or input_ids.shape[1] <= self.prompt_length: return scores generated = input_ids[:, self.prompt_length:] for batch_i in range(generated.shape[0]): used = torch.unique(generated[batch_i]) if used.numel() > 0: scores[batch_i, used] = scores[batch_i, used] - self.penalty return scores def decode_mol_token_ids(token_ids: torch.Tensor) -> str: """Decode a 1-D / 2-D tensor of generated mol-vocab IDs into a SMILES string. Mol generation runs in the decoder's own SMILES vocab (see :mod:`qwenvl.modalities.mol.tokenizer`); each id maps directly to a SMILES token via ``MOL_ID_TO_TOKEN``. ```` / ```` / ```` are dropped, and decoding stops at the first ````. """ from qwenvl.modalities.mol.tokenizer import MOL_ID_TO_TOKEN, MOL_SEP_ID if token_ids.dim() == 2: token_ids = token_ids[0] skip_tokens = {"", ""} out: List[str] = [] for tid in token_ids.tolist(): tid = int(tid) if tid == MOL_SEP_ID: break if 0 <= tid < len(MOL_ID_TO_TOKEN): tok = MOL_ID_TO_TOKEN[tid] if tok not in skip_tokens: out.append(tok) return "".join(out) def _checkpoint_key_names(model_path: str) -> Optional[set]: """Return checkpoint tensor keys when they can be inspected cheaply.""" root = Path(model_path) if not root.exists() or not root.is_dir(): return None for index_name in ("model.safetensors.index.json", "pytorch_model.bin.index.json"): index_path = root / index_name if index_path.exists(): try: with open(index_path, "r", encoding="utf-8") as f: data = json.load(f) weight_map = data.get("weight_map") or {} return set(weight_map.keys()) except Exception: return None safetensors_files = sorted(root.glob("*.safetensors")) if safetensors_files: try: from safetensors import safe_open except Exception: return None keys = set() try: for path in safetensors_files: with safe_open(str(path), framework="pt", device="cpu") as f: keys.update(f.keys()) return keys except Exception: return None return None _MOL_DECODER_REQUIRED_SUFFIXES = ( "tok_embed.weight", "kv_proj.0.weight", "layers.0.self_q.weight", "head.weight", ) _MOL_DECODER_VOCAB_SUFFIXES = ( "tok_embed.weight", "head.weight", ) def _current_mol_vocab_size() -> int: from qwenvl.modalities.mol.tokenizer import MOL_VOCAB_SIZE return int(MOL_VOCAB_SIZE) def _checkpoint_mol_config_vocab(model_path: str) -> Optional[int]: config_path = Path(model_path) / "config.json" if not config_path.exists(): return None with open(config_path, "r", encoding="utf-8") as f: data = json.load(f) mol_config = data.get("mol_config") if not isinstance(mol_config, dict): return None value = mol_config.get("mol_vocab_size") return int(value) if value is not None else None def _suffix_for_mol_decoder_key(key: str) -> Optional[str]: for suffix in _MOL_DECODER_REQUIRED_SUFFIXES: if key.endswith(f"decoders.mol.{suffix}"): return suffix return None def _checkpoint_mol_tensor_shapes( model_path: str, ) -> Optional[Dict[str, Tuple[str, Tuple[int, ...]]]]: """Return shapes for required mol decoder tensors, or None if unreadable.""" root = Path(model_path) if not root.exists() or not root.is_dir(): return None def _record(out, key, shape): suffix = _suffix_for_mol_decoder_key(key) if suffix is not None: out[suffix] = (key, tuple(int(x) for x in shape)) try: index_path = root / "model.safetensors.index.json" if index_path.exists(): from safetensors import safe_open with open(index_path, "r", encoding="utf-8") as f: weight_map = (json.load(f).get("weight_map") or {}) out = {} wanted_by_shard = {} for key, shard in weight_map.items(): if _suffix_for_mol_decoder_key(key) is not None: wanted_by_shard.setdefault(shard, []).append(key) for shard, keys in wanted_by_shard.items(): with safe_open(str(root / shard), framework="pt", device="cpu") as f: for key in keys: _record(out, key, f.get_tensor(key).shape) return out bin_index_path = root / "pytorch_model.bin.index.json" if bin_index_path.exists(): with open(bin_index_path, "r", encoding="utf-8") as f: weight_map = (json.load(f).get("weight_map") or {}) out = {} wanted_by_shard = {} for key, shard in weight_map.items(): if _suffix_for_mol_decoder_key(key) is not None: wanted_by_shard.setdefault(shard, []).append(key) for shard, keys in wanted_by_shard.items(): state = torch.load(str(root / shard), map_location="cpu") for key in keys: if key in state: _record(out, key, state[key].shape) return out safetensors_files = sorted(root.glob("*.safetensors")) if safetensors_files: from safetensors import safe_open out = {} for path in safetensors_files: with safe_open(str(path), framework="pt", device="cpu") as f: for key in f.keys(): if _suffix_for_mol_decoder_key(key) is not None: _record(out, key, f.get_tensor(key).shape) return out bin_files = sorted(root.glob("pytorch_model*.bin")) if bin_files: out = {} for path in bin_files: state = torch.load(str(path), map_location="cpu") for key, value in state.items(): if _suffix_for_mol_decoder_key(key) is not None: _record(out, key, value.shape) return out except Exception: return None return None def preflight_mol_checkpoint( model_path: str, fail_on_legacy_mol_decoder: bool = True, ) -> bool: """Strictly validate that a checkpoint uses the current mol vocab.""" expected_vocab = _current_mol_vocab_size() config_vocab = _checkpoint_mol_config_vocab(model_path) if config_vocab != expected_vocab: raise RuntimeError( "[mol-preflight] checkpoint mol_config.mol_vocab_size must match " f"current tokenizer vocab. checkpoint={config_vocab!r}, " f"expected={expected_vocab}. Re-train mol with the current tokenizer." ) keys = _checkpoint_key_names(model_path) if keys is None: raise RuntimeError( "[mol-preflight] could not inspect checkpoint tensor keys; refusing " "to load mol checkpoint because decoder vocab cannot be verified." ) missing = [ suffix for suffix in _MOL_DECODER_REQUIRED_SUFFIXES if not any(k.endswith(f"decoders.mol.{suffix}") for k in keys) ] if missing: has_any_mol_decoder = any("decoders.mol." in k for k in keys) raise RuntimeError( "[mol-preflight] checkpoint does not contain the required " f"MolARDecoder layout (missing: {', '.join(missing)}; " f"any_mol_decoder_keys={has_any_mol_decoder}). Re-train mol with " "the current code instead of loading a legacy/random decoder." ) shapes = _checkpoint_mol_tensor_shapes(model_path) if shapes is None: raise RuntimeError( "[mol-preflight] could not inspect mol decoder tensor shapes; refusing " "to load mol checkpoint because decoder vocab cannot be verified." ) missing_shapes = [s for s in _MOL_DECODER_VOCAB_SUFFIXES if s not in shapes] if missing_shapes: raise RuntimeError( "[mol-preflight] could not find mol decoder vocab tensor shapes for " f"{missing_shapes}; refusing to load this checkpoint." ) bad_shapes = [] for suffix in _MOL_DECODER_VOCAB_SUFFIXES: key, shape = shapes[suffix] if not shape or int(shape[0]) != expected_vocab: bad_shapes.append((key, shape)) if bad_shapes: detail = "; ".join(f"{key}: shape={shape}" for key, shape in bad_shapes) raise RuntimeError( "[mol-preflight] mol decoder tensor vocab dimension must match " f"current tokenizer vocab={expected_vocab}. {detail}. Re-train mol " "with the current tokenizer." ) print( "[mol-preflight] MolARDecoder vocab validation passed " f"(vocab={expected_vocab})." ) return True # --------------------------------------------------------------------------- # Helpers – mirror the training-time _build_messages logic # --------------------------------------------------------------------------- def _rna_input_placeholder(num_latent_tokens: int, placeholder_tag: str = "") -> str: """Placeholder for an INPUT RNA/DNA sequence (encoded by RNAConvFormer).""" tok = _RNA_TOKENS if placeholder_tag == "" else _DNA_TOKENS return tok["start"] + tok["pad"] * num_latent_tokens + tok["end"] def _protein_input_placeholder(num_latent_tokens: int) -> str: """Placeholder for an INPUT protein sequence (encoded by ESM3 + resampler).""" return _PROTEIN_TOKENS["start"] + _PROTEIN_TOKENS["pad"] * num_latent_tokens + _PROTEIN_TOKENS["end"] def _mol_input_placeholder(num_latent_tokens: int) -> str: """Placeholder for an INPUT molecular SMILES (encoded by GNN + resampler).""" return _MOL_TOKENS["start"] + _MOL_TOKENS["pad"] * num_latent_tokens + _MOL_TOKENS["end"] def build_inference_inputs( item: Dict[str, Any], processor, rna_tokenizer: Optional[RNACharTokenizer], protein_tokenizer=None, dna_tokenizer=None, num_latent_tokens: int = RNA_NUM_LATENT_TOKENS, num_dna_latent_tokens: int = RNA_NUM_LATENT_TOKENS, num_protein_latent_tokens: int = 32, num_mol_latent_tokens: int = 16, base_path: Path = Path(""), max_bio_seq_length: int = 4096, protein_max_residues: Optional[int] = None, rna_max_residues: Optional[int] = None, dna_max_residues: Optional[int] = None, mol_prompt_style: str = "prompt_only", mol_generation_slots: int = 0, ) -> Dict[str, torch.Tensor]: """ Convert a single sample dict (training format) into model-ready tensors. This mirrors preprocess_qwen_visual / _build_messages but for inference (no labels, only user turns matter, gpt turn is left for generation). """ if mol_prompt_style not in ("prompt_only", "train_slots"): raise ValueError( "mol_prompt_style must be one of {'prompt_only', 'train_slots'}, " f"got {mol_prompt_style!r}" ) mol_generation_slots = max(1, int(mol_generation_slots or 1)) # ---- extract media pools ---- rnas = item.get("rna") or [] if isinstance(rnas, str): rnas = [rnas] dnas = item.get("dna") or [] if isinstance(dnas, str): dnas = [dnas] proteins = item.get("protein") or [] if isinstance(proteins, str): proteins = [proteins] rna_cap = rna_max_residues if rna_max_residues is not None else max_bio_seq_length dna_cap = dna_max_residues if dna_max_residues is not None else max_bio_seq_length protein_cap = protein_max_residues if protein_max_residues is not None else max_bio_seq_length rnas = [s[: rna_cap - 2] for s in rnas] dnas = [s[: dna_cap - 2] for s in dnas] proteins = [s[: protein_cap - 2] for s in proteins] images = item.get("image") or [] if isinstance(images, str): images = [images] videos = item.get("video") or [] if isinstance(videos, str): videos = [videos] mols = item.get("mol") or [] if isinstance(mols, str): mols = [mols] # Hard guard: contact_map modality has been removed. if item.get("contact_map"): raise ValueError( "contact_map modality has been removed; drop the 'contact_map' " "field and any '' placeholders from your data." ) rna_pool = list(rnas) dna_pool = list(dnas) protein_pool = list(proteins) mol_pool = list(mols) image_pool = list(images) video_pool = list(videos) mol_smiles_list: List[str] = [] # SMILES consumed by placeholders # Sequences that need RNA encoder (input side only) input_sequences: List[str] = [] # DNA sequences (independent encoder; no longer aliased to RNA) dna_input_sequences: List[str] = [] # Protein sequences (separate encoder) protein_input_sequences: List[str] = [] image_processor = getattr(processor, "image_processor", None) has_train_slot_target = False # ---- replacement helpers (same logic as training) ---- def _replace_seq(text: str, tag: str, pool: list, is_assistant: bool) -> str: # Dispatches to the right input-sequences list based on the tag so # RNA and DNA flow through their own encoders. nonlocal input_sequences, dna_input_sequences # Replace only ONE occurrence — the outer _replace_all loop handles # interleaved ordering between different modality tags. if tag in text: if not pool: raise ValueError(f"More {tag} placeholders than sequences provided") seq = pool.pop(0) if is_assistant: text = text.replace(tag, "", 1) elif tag == "": dna_input_sequences.append(seq) text = text.replace(tag, _rna_input_placeholder(num_dna_latent_tokens, tag), 1) else: input_sequences.append(seq) text = text.replace(tag, _rna_input_placeholder(num_latent_tokens, tag), 1) return text def _replace_protein(text: str, is_assistant: bool) -> str: nonlocal protein_input_sequences # Replace only ONE occurrence — outer _replace_all controls ordering. if "" in text: if not protein_pool: raise ValueError("More placeholders than protein sequences provided") seq = protein_pool.pop(0) if is_assistant: text = text.replace("", "", 1) else: protein_input_sequences.append(seq) text = text.replace("", _protein_input_placeholder(num_protein_latent_tokens), 1) return text def _replace_mol( text: str, is_assistant: bool = False, is_last_assistant: bool = False, ) -> str: # Replace only ONE occurrence — outer _replace_all controls ordering. nonlocal has_train_slot_target if "" in text: is_mol_generation_target = ( is_assistant and is_last_assistant and mol_prompt_style == "train_slots" and (sample_task == "mol_generation" or text.strip() == "") ) if is_mol_generation_target: # Keep an assistant-side output window that matches the # training layout, but do not consume sample["mol"]. In eval # JSONL that field is the answer SMILES, so consuming it here # would leak the target into inference. has_train_slot_target = True return text.replace( "", BIO_SEQ_OUTPUT_PAD * mol_generation_slots, 1 ) if not mol_pool: raise ValueError("More placeholders than mol sequences provided") smiles = mol_pool.pop(0) if is_assistant: # Generation target: don't build graph, just remove placeholder. # The turn will be skipped anyway (last assistant turn), but we # consume from pool to keep ordering correct. text = text.replace("", "", 1) else: # Input: build graph for GNN encoder mol_smiles_list.append(smiles) text = text.replace("", _mol_input_placeholder(num_mol_latent_tokens), 1) return text def _replace_all( text: str, is_assistant: bool, is_last_assistant: bool = False, ) -> str: # Hard guard: contact_map modality has been removed. if "" in text: raise ValueError( "contact_map modality has been removed; remove " "'' placeholders from your data." ) while "" in text or "" in text or "" in text or "" in text: candidates = [] for tag in ("", "", "", ""): pos = text.find(tag) if pos >= 0: candidates.append((pos, tag)) if not candidates: break _, first_tag = min(candidates, key=lambda x: x[0]) if first_tag == "": text = _replace_seq(text, "", rna_pool, is_assistant) elif first_tag == "": text = _replace_seq(text, "", dna_pool, is_assistant) elif first_tag == "": text = _replace_protein(text, is_assistant) else: text = _replace_mol(text, is_assistant, is_last_assistant) return text # ---- build chat messages (only keep up to last user turn) ---- messages = [] conversations = item.get("conversations", []) # Find the index of the last assistant turn so we can skip it # (that's the response slot the model should generate). # Earlier assistant turns in multi-turn conversations are kept as context. last_assistant_idx = None for idx, turn in enumerate(conversations): if turn["from"] in ("gpt", "assistant"): last_assistant_idx = idx epi_cell_line = extract_epi_cell_line(item.get("task")) epi_injected = False sample_task = (item.get("task") or "").strip() # ``bool(dnas)`` checks the original list; dna_pool is mutated below. # Per-sample resolution: the caller may stamp ``_ea_label_mode`` on # each sample from --ea_label_mode, so the same checkpoint can be probed # in either mode without an env-var dance. _ea_n_bins = ea_n_bins_for_sample(item) if (bool(dnas) and sample_task == "EA") else None # Add standalone system prompt if present. system_prompt = item.get("system_prompt") or item.get("system") if system_prompt: text = system_prompt if _ea_n_bins is not None: text = rewrite_ea_binned_instruction(text, _ea_n_bins, "system") messages.append({"role": "system", "content": [{"type": "text", "text": text}]}) for turn_idx, turn in enumerate(conversations): role_raw = turn["from"] role = "user" if role_raw == "human" else ("system" if role_raw == "system" else "assistant") text: str = turn["value"] is_assistant = role == "assistant" if role == "user" and epi_cell_line and not epi_injected: text = f"{text} Cell line: {epi_cell_line}." epi_injected = True # Same EA instruction rewrite as training-time _build_messages, so the # model sees identical system/user instructions at inference. if role in ("system", "user") and _ea_n_bins is not None: text = rewrite_ea_binned_instruction(text, _ea_n_bins, role) is_last_assistant = is_assistant and turn_idx == last_assistant_idx text = _replace_all(text, is_assistant, is_last_assistant) if ( is_last_assistant and mol_prompt_style == "train_slots" and sample_task == "mol_generation" and not text.strip() ): # Some inference-formatted samples leave the final assistant turn # empty instead of using "". Still give MolARDecoder the # same kind of assistant-side hidden-state window it saw in # training. text = BIO_SEQ_OUTPUT_PAD * mol_generation_slots has_train_slot_target = True # Skip the last assistant turn — the model should generate this response. # This handles both empty turns (correct inference format) and # non-empty turns with ground truth (training format used for eval). if is_last_assistant and not ( mol_prompt_style == "train_slots" and BIO_SEQ_OUTPUT_PAD in text ): continue if role == "user" and ("" in text or "