ho22joshua's picture
Add agent-friendly inference and diagnostics
1375f74
Raw
History Blame Contribute Delete
5.14 kB
import os
import warnings
from functools import lru_cache
from threading import Thread
import gradio as gr
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from model_config import SYSTEM_PROMPT, resolve_model
from model_runtime import (
FAST_MAX_NEW_TOKENS,
FAST_TEMPERATURE,
FAST_TOP_P,
adapter_load_kwargs,
choose_device,
model_load_kwargs,
place_model,
)
MODEL_NAME = os.getenv("MODEL_NAME")
BASE_MODEL_ID = os.getenv("BASE_MODEL_ID")
ADAPTER_ID = os.getenv("ADAPTER_ID")
MODEL_MODE = os.getenv("MODEL_MODE", "adapter").lower()
ALLOW_CPU = os.getenv("ALLOW_CPU", "0") == "1"
DEVICE = os.getenv("DEVICE", "auto").lower()
FAST_MODE = os.getenv("FAST_MODE", "0") == "1"
@lru_cache(maxsize=1)
def load_model():
runtime = choose_device(ALLOW_CPU, DEVICE)
selected_model = resolve_model(MODEL_NAME, BASE_MODEL_ID, ADAPTER_ID)
base_model = selected_model["base_model"]
adapter = selected_model.get("adapter")
if MODEL_MODE == "adapter" and not adapter:
raise ValueError(
f"Model '{selected_model['name']}' does not define a LoRA adapter. "
"Set MODEL_MODE=base or choose a registry entry with an adapter."
)
if MODEL_MODE not in {"adapter", "base"}:
raise ValueError("MODEL_MODE must be either adapter or base.")
tokenizer_id = adapter if MODEL_MODE == "adapter" else base_model
tokenizer = AutoTokenizer.from_pretrained(tokenizer_id, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
base_model,
**model_load_kwargs(runtime),
)
if MODEL_MODE == "adapter":
model = PeftModel.from_pretrained(
model,
adapter,
**adapter_load_kwargs(runtime),
)
model = place_model(model, runtime)
model.eval()
return tokenizer, model
def respond(message, history, max_new_tokens, temperature, top_p):
tokenizer, model = load_model()
messages = [{"role": "system", "content": SYSTEM_PROMPT}]
for user_msg, assistant_msg in history:
if user_msg:
messages.append({"role": "user", "content": user_msg})
if assistant_msg:
messages.append({"role": "assistant", "content": assistant_msg})
messages.append({"role": "user", "content": message})
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
do_sample = float(temperature) > 0
if not do_sample and hasattr(model, "generation_config"):
model.generation_config.temperature = None
model.generation_config.top_p = None
model.generation_config.top_k = None
streamer = TextIteratorStreamer(
tokenizer,
skip_prompt=True,
skip_special_tokens=True,
)
generation_kwargs = {
**inputs,
"streamer": streamer,
"max_new_tokens": int(max_new_tokens),
"do_sample": do_sample,
"pad_token_id": tokenizer.eos_token_id,
}
if do_sample:
generation_kwargs["temperature"] = float(temperature)
generation_kwargs["top_p"] = float(top_p)
def generate_with_filtered_warnings():
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
message="To copy construct from a tensor.*",
category=UserWarning,
module="transformers.pytorch_utils",
)
model.generate(**generation_kwargs)
thread = Thread(target=generate_with_filtered_warnings)
thread.start()
partial = ""
for token in streamer:
partial += token
yield partial
with gr.Blocks(title="HEP Chat") as demo:
gr.Markdown(
"# HEP Chat\n"
"Ask the Qwen2.5-7B LoRA HEP assistant about signal signatures and Standard Model backgrounds."
)
gr.ChatInterface(
fn=respond,
additional_inputs=[
gr.Slider(
16,
1024,
value=FAST_MAX_NEW_TOKENS if FAST_MODE else 300,
step=1,
label="Max new tokens",
),
gr.Slider(
0.0,
1.5,
value=FAST_TEMPERATURE if FAST_MODE else 0.2,
step=0.05,
label="Temperature",
),
gr.Slider(
0.1,
1.0,
value=FAST_TOP_P if FAST_MODE else 0.9,
step=0.05,
label="Top-p",
),
],
examples=[
"For H to AA to photons, what Standard Model backgrounds should be considered?",
"For a final state with two leptons, b-jets, and missing transverse momentum, list dominant, irreducible, and reducible backgrounds.",
"What backgrounds mimic a diphoton plus missing transverse momentum search?",
],
)
if __name__ == "__main__":
demo.queue().launch(server_name=os.getenv("GRADIO_SERVER_NAME"))