contenteaseAI/LargeLanguageModel
0
1import gradio as gr2import os3import time4from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer, BitsAndBytesConfig5import torch6from threading import Thread7import logging8import spaces9from functools import lru_cache10 11# Set up logging12logging.basicConfig(level=logging.INFO)13logger = logging.getLogger(__name__)14 15# Set an environment variable16HF_TOKEN = os.environ.get("HF_TOKEN", None)17 18DESCRIPTION = '''19<div>20<h1 style="text-align: center;">ContenteaseAI custom trained model</h1>21</div>22'''23 24LICENSE = """25<p/>26---27For more information, visit our [website](https://contentease.ai).28"""29 30PLACEHOLDER = """31<div style="padding: 30px; text-align: center; display: flex; flex-direction: column; align-items: center;">32 <h1 style="font-size: 28px; margin-bottom: 2px; opacity: 0.55;">ContenteaseAI Custom AI trained model</h1>33 <p style="font-size: 18px; margin-bottom: 2px; opacity: 0.65;">Enter the text extracted from the PDF:</p>34</div>35"""36 37css = """38h1 {39 text-align: center;40 display: block;41}42"""43 44# Load the tokenizer and model with quantization45model_id = "meta-llama/Meta-Llama-3-8B-Instruct"46bnb_config = BitsAndBytesConfig(47 load_in_4bit=True,48 bnb_4bit_use_double_quant=True,49 bnb_4bit_quant_type="nf4",50 bnb_4bit_compute_dtype=torch.bfloat1651)52 53@lru_cache(maxsize=1)54def load_model_and_tokenizer():55 try:56 start_time = time.time()57 logger.info("Loading tokenizer...")58 tokenizer = AutoTokenizer.from_pretrained(model_id)59 logger.info("Loading model...")60 model = AutoModelForCausalLM.from_pretrained(61 model_id,62 device_map="auto",63 quantization_config=bnb_config,64 torch_dtype=torch.bfloat1665 )66 model.generation_config.pad_token_id = tokenizer.pad_token_id67 end_time = time.time()68 logger.info(f"Model and tokenizer loaded successfully in {end_time - start_time} seconds.")69 return model, tokenizer70 except Exception as e:71 logger.error(f"Error loading model or tokenizer: {e}")72 raise73 74try:75 model, tokenizer = load_model_and_tokenizer()76except Exception as e:77 logger.error(f"Failed to load model and tokenizer: {e}")78 raise79 80terminators = [81 tokenizer.eos_token_id,82 tokenizer.convert_tokens_to_ids("<|eot_id|>")83]84 85SYS_PROMPT = """86Extract all relevant keywords and add quantity from the following text and format the result in nested JSON, ignoring personal details and focusing only on the scope of work as shown in the example:87Good JSON example: {'lobby': {'frcm': {'replace': {'carpet': 1, 'carpet_pad': 1, 'base': 1, 'window_treatments': 1, 'artwork_and_decorative_accessories': 1, 'portable_lighting': 1, 'upholstered_furniture_and_decorative_pillows': 1, 'millwork': 1} } } }88Bad JSON example: {'lobby': { 'frcm': { 'replace': [ 'carpet', 'carpet_pad', 'base', 'window_treatments', 'artwork_and_decorative_accessories', 'portable_lighting', 'upholstered_furniture_and_decorative_pillows', 'millwork'] } } }89Make sure to fetch details from the provided text and ignore unnecessary information. The response should be in JSON format only, without any additional comments.90"""91 92def chunk_text(text, chunk_size=5000):93 """94 Splits the input text into chunks of specified size.95 96 Args:97 text (str): The input text to be chunked.98 chunk_size (int): The size of each chunk in tokens.99 100 Returns:101 list: A list of text chunks.102 """103 words = text.split()104 chunks = [' '.join(words[i:i + chunk_size]) for i in range(0, len(words), chunk_size)]105 return chunks106 107def combine_responses(responses):108 """109 Combines the responses from all chunks into a final output string.110 111 Args:112 responses (list): A list of responses from each chunk.113 114 Returns:115 str: The combined output string.116 """117 combined_output = " ".join(responses)118 return combined_output119 120def generate_response_for_chunk(chunk, history, temperature, max_new_tokens):121 start_time = time.time()122 123 conversation = [{"role": "system", "content": SYS_PROMPT}]124 for user, assistant in history:125 conversation.extend([{"role": "user", "content": user}, {"role": "assistant", "content": assistant}])126 conversation.append({"role": "user", "content": chunk})127 128 input_ids = tokenizer.apply_chat_template(conversation, return_tensors="pt").to(model.device)129 130 streamer = TextIteratorStreamer(tokenizer, timeout=10.0, skip_prompt=True, skip_special_tokens=True)131 132 generate_kwargs = dict(133 input_ids=input_ids,134 streamer=streamer,135 max_new_tokens=max_new_tokens,136 do_sample=True,137 temperature=temperature,138 eos_token_id=terminators,139 pad_token_id=tokenizer.eos_token_id140 )141 if temperature == 0:142 generate_kwargs['do_sample'] = False143 144 t = Thread(target=model.generate, kwargs=generate_kwargs)145 t.start()146 147 outputs = []148 for text in streamer:149 outputs.append(text)150 151 end_time = time.time()152 logger.info(f"Time taken for generating response for a chunk: {end_time - start_time} seconds")153 154 return "".join(outputs)155 156@spaces.GPU(duration=110)157def chat_llama3_8b(message: str, history: list, temperature: float, max_new_tokens: int):158 """159 Generate a streaming response using the llama3-8b model with chunking.160 161 Args:162 message (str): The input message.163 history (list): The conversation history used by ChatInterface.164 temperature (float): The temperature for generating the response.165 max_new_tokens (int): The maximum number of new tokens to generate.166 167 Returns:168 str: The generated response.169 """170 try:171 start_time = time.time()172 173 chunks = chunk_text(message)174 responses = []175 for chunk in chunks:176 response = generate_response_for_chunk(chunk, history, temperature, max_new_tokens)177 responses.append(response)178 final_output = combine_responses(responses)179 180 end_time = time.time()181 logger.info(f"Total time taken for generating response: {end_time - start_time} seconds")182 183 yield final_output184 except Exception as e:185 logger.error(f"Error generating response: {e}")186 yield "An error occurred while generating the response. Please try again."187 188# Gradio block189chatbot = gr.Chatbot(height=450, placeholder=PLACEHOLDER, label='Gradio ChatInterface')190 191with gr.Blocks(fill_height=True, css=css) as demo:192 gr.Markdown(DESCRIPTION)193 194 gr.ChatInterface(195 fn=chat_llama3_8b,196 chatbot=chatbot,197 fill_height=True,198 additional_inputs_accordion=gr.Accordion(label="⚙️ Parameters", open=False, render=False),199 additional_inputs=[200 gr.Slider(minimum=0, maximum=1, step=0.1, value=0.95, label="Temperature", render=False),201 gr.Slider(minimum=128, maximum=2000, step=1, value=700, label="Max new tokens", render=False),202 ]203 )204 205 gr.Markdown(LICENSE)206 207if __name__ == "__main__":208 try:209 demo.launch(show_error=True)210 except Exception as e:211 logger.error(f"Error launching Gradio demo: {e}")212 