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"))