github-actions[bot]
Sync inference Space from GitHub
79d1cc2
Raw
History Blame Contribute Delete
24.9 kB
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),
}