contenteaseAI/backup-LargeLanguageModel
0
1import gradio as gr2import os3import torch4from transformers import AutoConfig, AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer5from threading import Thread6from accelerate import init_empty_weights, infer_auto_device_map, disk_offload7 8# Set environment variables9HF_TOKEN = os.getenv("HF_TOKEN")10 11DESCRIPTION = '''12<div>13<h1 style="text-align: center;">ContenteaseAI custom trained model</h1>14</div>15'''16 17LICENSE = """18<p/>19 20---21For more information, visit our [website](https://contentease.ai).22"""23 24PLACEHOLDER = """25<div style="padding: 30px; text-align: center; display: flex; flex-direction: column; align-items: center;">26 <h1 style="font-size: 28px; margin-bottom: 2px; opacity: 0.55;">ContenteaseAI Custom AI trained model</h1>27 <p style="font-size: 18px; margin-bottom: 2px; opacity: 0.65;">Enter the text extracted from the PDF:</p>28</div>29"""30 31css = """32h1 {33 text-align: center;34 display: block;35}36"""37 38def initialize_model(model_name, max_memory=None):39 device = torch.device('cpu')40 41 # Load model configuration42 config = AutoConfig.from_pretrained(model_name)43 44 with init_empty_weights():45 # Initialize model with empty weights46 model = AutoModelForCausalLM.from_config(config)47 48 # Create device map based on memory constraints49 device_map = infer_auto_device_map(50 model, max_memory=max_memory, no_split_module_classes=["GPTNeoXLayer"], dtype="float16"51 )52 53 # Determine if offloading is needed54 needs_offloading = any(device == 'disk' for device in device_map.values())55 56 if needs_offloading:57 # Load model for offloading58 model = AutoModelForCausalLM.from_pretrained(59 model_name, device_map=device_map, offload_folder="offload",60 offload_state_dict=True, torch_dtype=torch.float1661 )62 offload_directory = "offload/"63 # Offload model to disk64 disk_offload(model=model, offload_dir=offload_directory)65 else:66 # Load model normally to specified device67 model = AutoModelForCausalLM.from_pretrained(68 model_name, torch_dtype=torch.float1669 )70 model.to(device)71 72 return model73 74try:75 # Initialize the model and tokenizer76 model_name = "meta-llama/Meta-Llama-3-8B-Instruct"77 model = initialize_model(model_name, max_memory={"cpu": "GiB"})78 tokenizer = AutoTokenizer.from_pretrained(model_name, use_auth_token=HF_TOKEN)79except Exception as e:80 print(f"Error initializing model: {e}")81 exit(1)82 83terminators = [84 tokenizer.eos_token_id,85 tokenizer.convert_tokens_to_ids("")86]87 88def chat_llama3_8b(message: str, history: list, temperature: float, max_new_tokens: int) -> str:89 """90 Generate a streaming response using the llama3-8b model.91 Args:92 message (str): The input message.93 history (list): The conversation history used by ChatInterface.94 temperature (float): The temperature for generating the response.95 max_new_tokens (int): The maximum number of new tokens to generate.96 Returns:97 str: The generated response.98 """99 conversation = []100 message += " Extract all relevant keywords and add quantity from the following text and format the result in nested JSON:"101 for user, assistant in history:102 conversation.extend([{"role": "user", "content": user}, {"role": "assistant", "content": assistant}])103 conversation.append({"role": "user", "content": message})104 105 input_ids = tokenizer.apply_chat_template(conversation, return_tensors="pt").to(model.device)106 107 streamer = TextIteratorStreamer(tokenizer, timeout=10.0, skip_prompt=True, skip_special_tokens=True)108 109 generate_kwargs = dict(110 input_ids=input_ids,111 streamer=streamer,112 max_new_tokens=max_new_tokens,113 do_sample=True,114 temperature=temperature,115 eos_token_id=terminators,116 )117 if temperature == 0:118 generate_kwargs['do_sample'] = False119 120 t = Thread(target=model.generate, kwargs=generate_kwargs)121 t.start()122 123 outputs = []124 for text in streamer:125 outputs.append(text)126 yield "".join(outputs)127 128# Gradio block129chatbot = gr.Chatbot(height=450, placeholder=PLACEHOLDER, label='Gradio ChatInterface')130 131with gr.Blocks(fill_height=True, css=css) as demo:132 gr.Markdown(DESCRIPTION)133 gr.ChatInterface(134 fn=chat_llama3_8b,135 chatbot=chatbot,136 fill_height=True,137 additional_inputs_accordion=gr.Accordion(label="⚙️ Parameters", open=False, render=False),138 additional_inputs=[139 gr.Slider(140 minimum=0,141 maximum=1,142 step=0.1,143 value=0.95,144 label="Temperature",145 render=False146 ),147 gr.Slider(148 minimum=128,149 maximum=9012,150 step=1,151 value=512,152 label="Max new tokens",153 render=False154 ),155 ]156 )157 gr.Markdown(LICENSE)158 159if __name__ == "__main__":160 demo.launch(server_port=8000, share=True)161 