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