ASL / app.py
Test502's picture
Update app.py
12bf817 verified
Raw
History Blame Contribute Delete
2.09 kB
import gradio as gr
import spaces
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
MODEL_ID = "saai-sa/ASL-4B-v1"
print("Loading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID,
trust_remote_code=True
)
model = None
@spaces.GPU(duration=120)
def chat(message, history):
global model
# تحميل المودل أول مرة فقط بعد الحصول على GPU
if model is None:
print("Loading ASL-4B on GPU...")
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
trust_remote_code=True
)
model = model.to("cuda")
model.eval()
print("ASL-4B loaded!")
messages = []
if history:
for item in history:
if isinstance(item, dict):
role = item.get("role")
content = item.get("content")
if role in ["user", "assistant"] and content:
messages.append({
"role": role,
"content": content
})
messages.append({
"role": "user",
"content": message
})
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
inputs = tokenizer(
prompt,
return_tensors="pt"
).to("cuda")
with torch.inference_mode():
output = model.generate(
**inputs,
max_new_tokens=512,
do_sample=True,
temperature=0.7,
top_p=0.8,
top_k=20,
repetition_penalty=1.0
)
generated_tokens = output[0][inputs.input_ids.shape[-1]:]
response = tokenizer.decode(
generated_tokens,
skip_special_tokens=True
)
return response
demo = gr.ChatInterface(
fn=chat,
title="ASL-4B-v1",
description="Saudi Arabic Language Model"
)
demo.queue()
demo.launch(
server_name="0.0.0.0",
server_port=7860,
ssr_mode=False
)