try: import tokenizers print(f"[inference] tokenizers import OK: {tokenizers.__version__}", flush=True) except Exception as e: print(f"[inference] tokenizers import FAILED: {e}", flush=True) import traceback traceback.print_exc() # Don't fail silently - this is the root cause of TokenizersBackend error raise import os import re import threading import hashlib from contextlib import nullcontext from importlib import metadata from typing import Literal import torch from fastapi import FastAPI, Header, HTTPException from pydantic import BaseModel from peft import PeftConfig, PeftModel, LoraConfig from transformers import ( AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, StoppingCriteria, StoppingCriteriaList, ) # --- SAFE PATCH for newer PEFT fields like alora_invocation_tokens --- if not hasattr(LoraConfig, '_is_patched'): _orig_lora_init = LoraConfig.__init__ def _patched_lora_init(self, *args, **kwargs): while True: try: return _orig_lora_init(self, *args, **kwargs) except TypeError as e: m = re.search(r"unexpected keyword argument '([^']+)'", str(e)) if not m: raise kwargs.pop(m.group(1), None) LoraConfig.__init__ = _patched_lora_init LoraConfig._is_patched = True # FIXED TEMPLATE - no nested if/else confusion MISTRAL_STAGE_CHAT_TEMPLATE = ( "{% if messages[0]['role'] == 'system' %}" "{% set system_message = messages[0]['content'] %}" "{% set messages = messages[1:] %}" "{% else %}" "{% set system_message = '' %}" "{% endif %}" "{% for message in messages %}" "{% if message['role'] == 'user' %}" "{% if loop.first and system_message != '' %}" "{{ bos_token + '[INST] ' + system_message.strip() + '\\n\\n' + message['content'].strip() + ' [/INST]' }}" "{% else %}" "{{ bos_token + '[INST] ' + message['content'].strip() + ' [/INST]' }}" "{% endif %}" "{% elif message['role'] == 'assistant' %}" "{{ ' ' + message['content'].strip() + eos_token }}" "{% endif %}" "{% endfor %}" ) ADAPTER_MODEL_ID = os.environ.get( "ADAPTER_MODEL_ID", "maimd/Maimd-HPI-SFT-Behavioral-BioMistral-7B-v001-20260609", ).strip() MODEL_ID = os.environ.get("MODEL_ID", "").strip() BASE_MODEL_ID = os.environ.get("BASE_MODEL_ID", "").strip() TOKENIZER_MODEL_ID = os.environ.get("TOKENIZER_MODEL_ID", "").strip() HF_TOKEN = os.environ.get("HF_TOKEN") REMOTE_API_KEY = os.environ.get("REMOTE_API_KEY", "") HPI_ADAPTER_ID = os.environ.get("HPI_ADAPTER_ID", "").strip() if not HPI_ADAPTER_ID and not BASE_MODEL_ID: HPI_ADAPTER_ID = ADAPTER_MODEL_ID ADAPTER_IDS = { "hpi": HPI_ADAPTER_ID, "pe": os.environ.get("PE_ADAPTER_ID", "").strip(), "ddx": os.environ.get("DDX_ADAPTER_ID", "").strip(), "cdf": os.environ.get("CDF_ADAPTER_ID", "").strip(), } app = FastAPI(title="Agentic Dr Inference") model = None tokenizer = None resolved_model_id = None resolved_base_model_id = None resolved_tokenizer_source = None loaded_adapters = set() active_adapter = None model_load_error = None inference_lock = threading.RLock() class ChatMessage(BaseModel): role: Literal["system", "user", "assistant"] content: str class GenerateRequest(BaseModel): prompt: str = "" max_new_tokens: int = 300 adapter_key: str = "hpi" response_format: Literal["auto", "text", "json"] = "auto" messages: list[ChatMessage] | None = None do_sample: bool = False temperature: float = 0.0 top_p: float = 1.0 seed: int = 42 def _complete_json_end(text: str) -> int | None: """Return the end offset of the first complete top-level JSON object.""" start = text.find("{") if start < 0: return None depth = 0 in_string = False escaped = False for index, char in enumerate(text[start:], start=start): if escaped: escaped = False continue if char == "\\" and in_string: escaped = True continue if char == '"': in_string = not in_string continue if in_string: continue if char == "{": depth += 1 elif char == "}": depth -= 1 if depth == 0: return index + 1 return None class CompleteJsonStoppingCriteria(StoppingCriteria): """Stop structured stages immediately after one complete JSON object.""" def __init__(self, tokenizer_instance, prompt_length: int) -> None: self.tokenizer = tokenizer_instance self.prompt_length = prompt_length def __call__(self, input_ids, scores, **kwargs) -> bool: generated = input_ids[0][self.prompt_length:] text = self.tokenizer.decode(generated, skip_special_tokens=True) return _complete_json_end(text) is not None def _minimum_structured_generation_tokens(adapter_key: str) -> int: """Minimum EOS-free window for each reviewed structured stage.""" return { "hpi": 300, "pe": 200, "ddx": 350, "cdf": 300, }.get(adapter_key.strip().lower(), 64) def log_inference_event(event: str, **kwargs): print(f"[inference] {event} | {kwargs}", flush=True) def adapter_id_for_stage(adapter_key: str) -> str: normalized_key = adapter_key.strip().lower() adapter_id = ADAPTER_IDS.get(normalized_key, "").strip() if not adapter_id: raise HTTPException( status_code=400, detail=( f"No LoRA adapter configured for stage '{normalized_key}'. " f"Set {normalized_key.upper()}_ADAPTER_ID in the inference Space." ), ) return adapter_id def _normalized_model_id(model_id: str | None) -> str: return str(model_id or "").strip().rstrip("/").lower() def _adapter_parent_model_id(adapter_key: str, adapter_id: str) -> str: try: config = PeftConfig.from_pretrained(adapter_id, token=HF_TOKEN) except Exception as exc: raise RuntimeError( f"Could not read the {adapter_key} adapter config from {adapter_id}: {exc}" ) from exc return str(config.base_model_name_or_path or "").strip() def _require_compatible_adapter_parent( adapter_key: str, adapter_id: str, expected_base_model_id: str, ) -> None: parent_model_id = _adapter_parent_model_id(adapter_key, adapter_id) if _normalized_model_id(parent_model_id) != _normalized_model_id(expected_base_model_id): raise RuntimeError( f"{adapter_key} adapter/base mismatch: {adapter_id} was trained against " f"{parent_model_id}, but the inference resident base is " f"{expected_base_model_id}. LoRA adapters must run on the checkpoint " "they were trained against." ) def _tokenizer_source_candidates( *, base_model_id: str, adapter_model_id: str | None, target_model_id: str, ) -> list[str]: candidates = [ TOKENIZER_MODEL_ID, base_model_id, adapter_model_id or "", target_model_id, ] if "biomistral" in " ".join(candidates).lower(): candidates.append("BioMistral/BioMistral-7B") unique_sources: list[str] = [] for candidate in candidates: source = str(candidate or "").strip() if source and source not in unique_sources: unique_sources.append(source) return unique_sources def _load_compatible_tokenizer( *, base_model_id: str, adapter_model_id: str | None, target_model_id: str, ): failures: list[str] = [] sources = _tokenizer_source_candidates( base_model_id=base_model_id, adapter_model_id=adapter_model_id, target_model_id=target_model_id, ) for source in sources: for use_fast in (True, False): try: log_inference_event( "tokenizer_loading", tokenizer_source=source, use_fast=use_fast, ) loaded = AutoTokenizer.from_pretrained( source, token=HF_TOKEN, use_fast=use_fast, trust_remote_code=True, ) log_inference_event( "tokenizer_load_succeeded", tokenizer_source=source, use_fast=use_fast, tokenizer_class=type(loaded).__name__, ) return loaded, source except Exception as exc: error = f"{type(exc).__name__}: {exc}" failures.append(f"{source} use_fast={use_fast}: {error}") log_inference_event( "tokenizer_load_failed", tokenizer_source=source, use_fast=use_fast, error=error, ) failure_summary = " | ".join(failures[-6:]) raise RuntimeError( "No compatible tokenizer could be loaded. Set TOKENIZER_MODEL_ID to the " f"tokenizer repository used during training. Attempts: {failure_summary}" ) def package_versions(): versions = {} for package_name in ("torch", "transformers", "tokenizers", "accelerate", "bitsandbytes", "peft"): try: versions[package_name] = metadata.version(package_name) except metadata.PackageNotFoundError: versions[package_name] = "not_installed" return versions def load_model(): global model, tokenizer, resolved_tokenizer_source, active_adapter, model_load_error if model is not None and tokenizer is not None: return with inference_lock: if model is not None and tokenizer is not None: return try: _load_model_unlocked() model_load_error = None except Exception as exc: model = None tokenizer = None resolved_tokenizer_source = None active_adapter = None loaded_adapters.clear() model_load_error = f"{type(exc).__name__}: {exc}" log_inference_event( "model_load_failed", error=model_load_error, package_versions=package_versions(), ) if torch.cuda.is_available(): torch.cuda.empty_cache() raise def _load_model_unlocked(): global model, tokenizer, resolved_model_id, resolved_base_model_id global resolved_tokenizer_source, active_adapter log_inference_event( "adapter_configuration", base_model_id_configured=bool(BASE_MODEL_ID), model_id_configured=bool(MODEL_ID), tokenizer_model_id=TOKENIZER_MODEL_ID, hpi_generation_route="adapter" if ADAPTER_IDS["hpi"] else "base_model", adapter_ids=ADAPTER_IDS, configured_adapters={key: bool(value) for key, value in ADAPTER_IDS.items()}, package_versions=package_versions(), ) device = "cuda" if torch.cuda.is_available() else "cpu" quantization_config = None if device == "cuda": quantization_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", ) target_model_id = BASE_MODEL_ID or MODEL_ID or ADAPTER_IDS.get("hpi") or ADAPTER_MODEL_ID if not target_model_id: raise RuntimeError("At least one of HPI_ADAPTER_ID, ADAPTER_MODEL_ID, MODEL_ID, or BASE_MODEL_ID must be configured.") resolved_model_id = target_model_id adapter_model_id = None if BASE_MODEL_ID: resolved_base_model_id = BASE_MODEL_ID if ADAPTER_IDS.get("hpi"): try: peft_config = PeftConfig.from_pretrained( ADAPTER_IDS["hpi"], token=HF_TOKEN, ) adapter_model_id = ADAPTER_IDS["hpi"] if not resolved_base_model_id: resolved_base_model_id = peft_config.base_model_name_or_path except Exception as exc: log_inference_event( "peft_config_load_failed", adapter_key="hpi", adapter_id=ADAPTER_IDS["hpi"], error=f"{type(exc).__name__}: {exc}", base_model_id_configured=bool(BASE_MODEL_ID), has_hf_token=bool(HF_TOKEN), ) if not BASE_MODEL_ID: raise RuntimeError(f"Could not resolve base model from HPI adapter: {exc}") from exc adapter_model_id = ADAPTER_IDS["hpi"] if not resolved_base_model_id: resolved_base_model_id = BASE_MODEL_ID if not resolved_base_model_id: resolved_base_model_id = target_model_id for adapter_key, configured_adapter_id in ADAPTER_IDS.items(): if configured_adapter_id: _require_compatible_adapter_parent( adapter_key, configured_adapter_id, resolved_base_model_id, ) tokenizer, resolved_tokenizer_source = _load_compatible_tokenizer( base_model_id=resolved_base_model_id, adapter_model_id=adapter_model_id, target_model_id=target_model_id, ) if "mistral" in str(resolved_base_model_id or target_model_id).lower(): tokenizer.chat_template = MISTRAL_STAGE_CHAT_TEMPLATE if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token log_inference_event( "tokenizer_ready", tokenizer_source=resolved_tokenizer_source, tokenizer_class=type(tokenizer).__name__, tokenizer_is_fast=bool(getattr(tokenizer, "is_fast", False)), custom_mistral_chat_template=(tokenizer.chat_template == MISTRAL_STAGE_CHAT_TEMPLATE), ) base_model = AutoModelForCausalLM.from_pretrained( resolved_base_model_id or target_model_id, token=HF_TOKEN, torch_dtype=torch.bfloat16 if device == "cuda" else torch.float32, device_map="auto" if device == "cuda" else None, low_cpu_mem_usage=True, quantization_config=quantization_config, ) configured_adapters = [ (key, adapter_id) for key, adapter_id in ADAPTER_IDS.items() if adapter_id ] if configured_adapters: first_adapter_key, first_adapter_id = configured_adapters[0] model = PeftModel.from_pretrained( base_model, first_adapter_id, adapter_name=first_adapter_key, token=HF_TOKEN, ).eval() loaded_adapters.add(first_adapter_key) active_adapter = first_adapter_key log_inference_event( "adapter_loaded", adapter_key=first_adapter_key, adapter_id=first_adapter_id, base_model_id=resolved_base_model_id, loaded_adapters=sorted(loaded_adapters), active_adapter=active_adapter, ) else: model = base_model for k, aid in configured_adapters: if not aid or k in loaded_adapters: continue try: log_inference_event("adapter_preloading", adapter_key=k, adapter_id=aid) model.load_adapter(aid, adapter_name=k, token=HF_TOKEN) loaded_adapters.add(k) log_inference_event("adapter_preloaded", adapter_key=k, loaded_adapters=sorted(loaded_adapters)) except Exception as e: log_inference_event("adapter_preload_failed", adapter_key=k, error=f"{type(e).__name__}: {e}") model.eval() log_inference_event( "model_ready", base_model_id=resolved_base_model_id, active_adapter=active_adapter, loaded_adapters=sorted(loaded_adapters), ) def activate_adapter(adapter_key: str) -> str: global active_adapter load_model() normalized_key = adapter_key.strip().lower() if not normalized_key: normalized_key = "hpi" if normalized_key == "hpi" and not ADAPTER_IDS["hpi"]: if not BASE_MODEL_ID: raise HTTPException( status_code=500, detail="HPI base-model routing requires BASE_MODEL_ID.", ) return "base" adapter_id = adapter_id_for_stage(normalized_key) if not isinstance(model, PeftModel): raise HTTPException( status_code=500, detail="The inference model was loaded without PEFT support, so LoRA adapters cannot be switched.", ) if normalized_key not in loaded_adapters: log_inference_event( "adapter_loading", adapter_key=normalized_key, adapter_id=adapter_id, active_adapter_before=active_adapter, loaded_adapters=sorted(loaded_adapters), ) model.load_adapter( adapter_id, adapter_name=normalized_key, token=HF_TOKEN, ) loaded_adapters.add(normalized_key) log_inference_event( "adapter_loaded", adapter_key=normalized_key, adapter_id=adapter_id, active_adapter=active_adapter, loaded_adapters=sorted(loaded_adapters), ) if active_adapter != normalized_key: previous_adapter = active_adapter model.set_adapter(normalized_key) active_adapter = normalized_key log_inference_event( "adapter_activated", previous_adapter=previous_adapter, active_adapter=active_adapter, adapter_key=normalized_key, adapter_id=adapter_id, loaded_adapters=sorted(loaded_adapters), ) return normalized_key @app.get("/health") def health(): if model is not None: model_status = "ready" elif model_load_error: model_status = "load_failed" else: model_status = "not_loaded" return { "status": "ok", "model_status": model_status, "model_id": resolved_model_id or MODEL_ID or ADAPTER_MODEL_ID, "base_model_id": resolved_base_model_id, "tokenizer_source": resolved_tokenizer_source, "tokenizer_is_fast": bool(getattr(tokenizer, "is_fast", False)) if tokenizer else None, "configured_adapters": {key: bool(value) for key, value in ADAPTER_IDS.items()}, "hpi_generation_route": "adapter" if ADAPTER_IDS["hpi"] else "base", "loaded_adapters": sorted(loaded_adapters), "active_adapter": active_adapter, "model_load_error": model_load_error, "package_versions": package_versions(), } @app.post("/generate") def generate(request: GenerateRequest, authorization: str | None = Header(default=None)): if REMOTE_API_KEY: expected = f"Bearer {REMOTE_API_KEY}" if authorization != expected: raise HTTPException(status_code=401, detail="unauthorized") if not request.prompt.strip() and not request.messages: raise HTTPException(status_code=400, detail="prompt or messages is required") with inference_lock: try: generation_route = activate_adapter(request.adapter_key) log_inference_event( "generation_route_selected", adapter_key=request.adapter_key, generation_route=generation_route, base_model_id=resolved_base_model_id, configured_adapter_id=ADAPTER_IDS.get( request.adapter_key.strip().lower(), "", ), ) messages = ( [ {"role": message.role, "content": message.content} for message in request.messages ] if request.messages else [{"role": "user", "content": request.prompt}] ) prompt_text = tokenizer.apply_chat_template( messages, add_generation_prompt=True, tokenize=False, ) inputs = tokenizer(prompt_text, return_tensors="pt") model_input_device = next(model.parameters()).device inputs = { key: value.to(model_input_device) if isinstance(value, torch.Tensor) else value for key, value in inputs.items() } input_len = inputs["input_ids"].shape[-1] normalized_adapter = request.adapter_key.strip().lower() if request.response_format == "json": expects_json = True elif request.response_format == "text": expects_json = False else: expects_json = ( normalized_adapter != "hpi" or request.max_new_tokens > 200 ) torch.manual_seed(request.seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(request.seed) gen_kwargs = { "max_new_tokens": request.max_new_tokens, "do_sample": request.do_sample, "pad_token_id": tokenizer.eos_token_id, "eos_token_id": tokenizer.eos_token_id, "use_cache": True, } if request.do_sample: gen_kwargs.update( { "temperature": max(request.temperature, 1e-5), "top_p": min(max(request.top_p, 1e-5), 1.0), } ) if expects_json: # Prevent a structured generation from emitting one opening brace # followed immediately by EOS. Complete JSON can still stop earlier # through CompleteJsonStoppingCriteria. gen_kwargs["min_new_tokens"] = min( _minimum_structured_generation_tokens(normalized_adapter), request.max_new_tokens, ) gen_kwargs["stopping_criteria"] = StoppingCriteriaList( [CompleteJsonStoppingCriteria(tokenizer, input_len)] ) adapter_context = ( model.disable_adapter() if generation_route == "base" and isinstance(model, PeftModel) else nullcontext() ) with adapter_context, torch.inference_mode(): outputs = model.generate(**inputs, **gen_kwargs) generated_tokens = outputs[0][input_len:] text = tokenizer.decode(generated_tokens, skip_special_tokens=True).strip() json_end = _complete_json_end(text) if expects_json else None if json_end is not None: text = text[:json_end].strip() elif expects_json: output_sha256 = hashlib.sha256(text.encode("utf-8")).hexdigest() log_inference_event( "structured_generation_incomplete", adapter_key=request.adapter_key, generated_character_count=len(text), generated_token_count=len(generated_tokens), output_sha256=output_sha256, generated_text_preview=text[-800:], ) raise HTTPException( status_code=502, detail=( "Structured generation ended before a complete JSON object " f"was produced; output_sha256={output_sha256}" ), ) except HTTPException: raise except Exception as exc: error = f"{type(exc).__name__}: {exc}" log_inference_event( "generation_failed", adapter_key=request.adapter_key, active_adapter=active_adapter, loaded_adapters=sorted(loaded_adapters), error=error, ) raise HTTPException( status_code=500, detail=f"Inference generation failed: {error}", ) from exc log_inference_event( "generation_completed", adapter_key=request.adapter_key, active_adapter=active_adapter, generation_route=generation_route, generated_character_count=len(text), generated_token_count=len(generated_tokens), deterministic=not request.do_sample, input_sha256=hashlib.sha256(prompt_text.encode("utf-8")).hexdigest(), generated_text_preview=text[-800:], ) return { "text": text, "adapter_key": request.adapter_key, "deterministic": not request.do_sample, "generated_tokens": len(generated_tokens), }