import os from collections.abc import Iterator from threading import Thread import gradio as gr import spaces import torch from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, TextIteratorStreamer MAX_NEW_TOKENS_LIMIT = 2048 DEFAULT_MAX_NEW_TOKENS = 1024 MAX_INPUT_TOKENS = int(os.getenv("MAX_INPUT_TOKENS", "4096")) DESCRIPTION = """\ # Llama-2 7B Chat This Space demonstrates model [Llama-2-7b-chat](https://huggingface.co/meta-llama/Llama-2-7b-chat) by Meta, a Llama 2 model with 7B parameters fine-tuned for chat instructions. Feel free to play with it, or duplicate to run generations without a queue! If you want to run your own service, you can also [deploy the model on Inference Endpoints](https://huggingface.co/inference-endpoints). 🔎 For more details about the Llama 2 family of models and how to use them with `transformers`, take a look [at our blog post](https://huggingface.co/blog/llama2). 🔨 Looking for an even more powerful model? Check out the [13B version](https://huggingface.co/spaces/huggingface-projects/llama-2-13b-chat) or the large [70B model demo](https://huggingface.co/spaces/ysharma/Explore_llamav2_with_TGI). """ LICENSE_NOTICE = """

--- As a derivate work of [Llama-2-7b-chat](https://huggingface.co/meta-llama/Llama-2-7b-chat) by Meta, this demo is governed by the original [license](https://huggingface.co/spaces/huggingface-projects/llama-2-7b-chat/blob/main/LICENSE.txt) and [acceptable use policy](https://huggingface.co/spaces/huggingface-projects/llama-2-7b-chat/blob/main/USE_POLICY.md). """ MODEL_ID = "meta-llama/Llama-2-7b-chat-hf" model = AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.float16, device_map="auto") tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) tokenizer.use_default_system_prompt = False class StopOnSignal(StoppingCriteria): def __init__(self) -> None: self.stopped = False def __call__(self, input_ids: torch.Tensor, scores: torch.Tensor, **kwargs: object) -> bool: # noqa: ARG002 return self.stopped @spaces.GPU def _generate_on_gpu( input_ids: torch.Tensor, max_new_tokens: int, temperature: float, top_p: float, top_k: int, repetition_penalty: float, ) -> Iterator[str]: input_ids = input_ids.to(model.device) streamer = TextIteratorStreamer(tokenizer, timeout=20.0, skip_prompt=True, skip_special_tokens=True) stop_criteria = StopOnSignal() generate_kwargs = { "input_ids": input_ids, "streamer": streamer, "stopping_criteria": [stop_criteria], "max_new_tokens": max_new_tokens, "max_length": None, "do_sample": True, "top_p": top_p, "top_k": top_k, "temperature": temperature, "num_beams": 1, "repetition_penalty": repetition_penalty, } exception_holder: list[Exception] = [] def _generate() -> None: try: model.generate(**generate_kwargs) except Exception as e: # noqa: BLE001 exception_holder.append(e) thread = Thread(target=_generate) thread.start() chunks: list[str] = [] try: for text in streamer: chunks.append(text) yield "".join(chunks) except GeneratorExit: stop_criteria.stopped = True for _ in streamer: pass thread.join() raise thread.join() if exception_holder: error_msg = f"Generation failed: {exception_holder[0]}" raise gr.Error(error_msg) def validate_input(message: str) -> dict: return gr.validate(bool(message and message.strip()), "Please enter a message.") def generate( message: str, chat_history: list[dict], system_prompt: str = "", max_new_tokens: int = 1024, temperature: float = 0.6, top_p: float = 0.9, top_k: int = 50, repetition_penalty: float = 1.2, ) -> Iterator[str]: conversation = [] if system_prompt: conversation.append({"role": "system", "content": system_prompt}) for msg in chat_history: if isinstance(msg["content"], list): text = "".join(part["text"] for part in msg["content"] if part["type"] == "text") else: text = str(msg["content"]) conversation.append({"role": msg["role"], "content": text}) conversation.append({"role": "user", "content": message}) input_ids = tokenizer.apply_chat_template(conversation, return_tensors="pt", return_dict=True).input_ids n_input_tokens = input_ids.shape[1] if n_input_tokens > MAX_INPUT_TOKENS: error_msg = f"Input too long ({n_input_tokens} tokens). Maximum is {MAX_INPUT_TOKENS} tokens." raise gr.Error(error_msg) max_new_tokens = min(max_new_tokens, MAX_INPUT_TOKENS - n_input_tokens) if max_new_tokens <= 0: raise gr.Error("Input uses the entire context window. No room to generate new tokens.") yield from _generate_on_gpu( input_ids=input_ids, max_new_tokens=max_new_tokens, temperature=temperature, top_p=top_p, top_k=top_k, repetition_penalty=repetition_penalty, ) chat_interface = gr.ChatInterface( fn=generate, validator=validate_input, chatbot=gr.Chatbot(height="70vh"), additional_inputs=[ gr.Textbox(label="System prompt", lines=6), gr.Slider( label="Max new tokens", minimum=1, maximum=MAX_NEW_TOKENS_LIMIT, step=1, value=DEFAULT_MAX_NEW_TOKENS, ), gr.Slider( label="Temperature", minimum=0.1, maximum=4.0, step=0.1, value=0.6, ), gr.Slider( label="Top-p (nucleus sampling)", minimum=0.05, maximum=1.0, step=0.05, value=0.9, ), gr.Slider( label="Top-k", minimum=1, maximum=1000, step=1, value=50, ), gr.Slider( label="Repetition penalty", minimum=1.0, maximum=2.0, step=0.05, value=1.2, ), ], examples=[ ["Hello there! How are you doing?"], ["Can you explain briefly to me what is the Python programming language?"], ["Explain the plot of Cinderella in a sentence."], ["How many hours does it take a man to eat a Helicopter?"], ["Write a 100-word article on 'Benefits of Open-Source in AI research'"], ], cache_examples=False, ) with gr.Blocks() as demo: gr.Markdown(DESCRIPTION) chat_interface.render() gr.Markdown(LICENSE_NOTICE) if __name__ == "__main__": demo.launch(css_paths="style.css")