Team Ai
Apppublic

contenteaseAI/backup-LargeLanguageModel

sourceHugging Facellama3updated 2y agoView on Hugging Face
0likes
utils.py161 linesDownload Raw Back to root
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