| import gradio as gr |
| import torch |
|
|
| from transformers import ( |
| AutoTokenizer, |
| AutoModelForCausalLM |
| ) |
|
|
|
|
| MODEL_ID = "jerinaj/lfm-tool-merged" |
|
|
|
|
| |
| tokenizer = AutoTokenizer.from_pretrained( |
| MODEL_ID |
| ) |
|
|
|
|
| |
| model = AutoModelForCausalLM.from_pretrained( |
| MODEL_ID, |
| dtype=torch.bfloat16, |
| device_map="auto" |
| ) |
|
|
| model.eval() |
|
|
|
|
| def predict(messages): |
|
|
| prompt = tokenizer.apply_chat_template( |
| messages, |
| tokenize=False, |
| add_generation_prompt=True |
| ) |
|
|
| inputs = tokenizer( |
| prompt, |
| return_tensors="pt" |
| ) |
|
|
| inputs = { |
| k: v.to(model.device) |
| for k, v in inputs.items() |
| } |
|
|
|
|
| with torch.inference_mode(): |
|
|
| output = model.generate( |
| **inputs, |
| max_new_tokens=256, |
| temperature=0.1, |
| do_sample=True |
| ) |
|
|
|
|
| response = tokenizer.decode( |
| output[0][inputs["input_ids"].shape[1]:], |
| skip_special_tokens=True |
| ) |
|
|
| return response |
|
|
|
|
|
|
| demo = gr.Interface( |
| fn=predict, |
| inputs=gr.JSON(), |
| outputs=gr.Textbox() |
| ) |
|
|
|
|
| demo.launch( |
| server_name="0.0.0.0", |
| server_port=7860, |
| share=True |
| ) |