| import os |
| import gradio as gr |
| import numpy as np |
| import onnxruntime as ort |
| import sentencepiece as spm |
| from huggingface_hub import snapshot_download |
|
|
| MODEL_REPO = "egomnia/emma-5" |
| MAX_CONTEXT = 2048 |
|
|
| model_dir = snapshot_download( |
| repo_id=MODEL_REPO, |
| allow_patterns=[ |
| "*.onnx", |
| "*.onnx.data", |
| "*.model", |
| "*.json" |
| ] |
| ) |
|
|
| files = os.listdir(model_dir) |
|
|
| onnx_file = next( |
| os.path.join(model_dir, f) |
| for f in files |
| if f.endswith(".onnx") |
| ) |
|
|
| tokenizer_file = next( |
| os.path.join(model_dir, f) |
| for f in files |
| if f.endswith(".model") |
| ) |
|
|
| sp = spm.SentencePieceProcessor(model_file=tokenizer_file) |
|
|
| session = ort.InferenceSession( |
| onnx_file, |
| providers=["CPUExecutionProvider"] |
| ) |
|
|
| eos_id = sp.eos_id() |
|
|
|
|
| def sample_token(logits, temperature, top_p): |
| logits = logits.astype(np.float64) |
|
|
| if temperature <= 0: |
| return int(np.argmax(logits)) |
|
|
| logits = logits / temperature |
|
|
| logits = logits - np.max(logits) |
| probs = np.exp(logits) |
| probs = probs / np.sum(probs) |
|
|
| sorted_ids = np.argsort(probs)[::-1] |
| sorted_probs = probs[sorted_ids] |
|
|
| cumulative_probs = np.cumsum(sorted_probs) |
|
|
| cutoff = cumulative_probs > top_p |
|
|
| if np.any(cutoff): |
| first_cutoff = np.argmax(cutoff) |
| sorted_probs[first_cutoff + 1:] = 0 |
|
|
| sorted_probs = sorted_probs / np.sum(sorted_probs) |
|
|
| selected_id = np.random.choice( |
| sorted_ids, |
| p=sorted_probs |
| ) |
|
|
| return int(selected_id) |
|
|
|
|
| def generate(prompt, max_new_tokens, temperature, top_p): |
| if not prompt or not prompt.strip(): |
| return "Scrivi un prompt." |
|
|
| token_ids = sp.encode(prompt, out_type=int) |
| token_ids = token_ids[-MAX_CONTEXT:] |
|
|
| generated_ids = [] |
|
|
| for _ in range(int(max_new_tokens)): |
| current_ids = token_ids[-MAX_CONTEXT:] |
|
|
| input_ids = np.array( |
| [current_ids], |
| dtype=np.int64 |
| ) |
|
|
| logits = session.run( |
| None, |
| {"input_ids": input_ids} |
| )[0] |
|
|
| next_token_logits = logits[0, -1, :] |
|
|
| next_token_id = sample_token( |
| next_token_logits, |
| temperature, |
| top_p |
| ) |
|
|
| if next_token_id == eos_id: |
| break |
|
|
| token_ids.append(next_token_id) |
| generated_ids.append(next_token_id) |
|
|
| generated_text = sp.decode(generated_ids) |
|
|
| return prompt + generated_text |
|
|
|
|
| with gr.Blocks(title="Emma-5 Playground") as demo: |
| gr.Markdown( |
| """ |
| # Emma-5 Playground |
| |
| Mini LLM italiano Emma-5 eseguito in CPU tramite ONNX Runtime. |
| """ |
| ) |
|
|
| prompt = gr.Textbox( |
| label="Prompt", |
| placeholder="Scrivi qualcosa...", |
| lines=5 |
| ) |
|
|
| with gr.Row(): |
| max_tokens = gr.Slider( |
| minimum=1, |
| maximum=150, |
| value=50, |
| step=1, |
| label="Token da generare" |
| ) |
|
|
| temperature = gr.Slider( |
| minimum=0, |
| maximum=2, |
| value=0.8, |
| step=0.05, |
| label="Temperature" |
| ) |
|
|
| top_p = gr.Slider( |
| minimum=0.1, |
| maximum=1, |
| value=0.9, |
| step=0.05, |
| label="Top-p" |
| ) |
|
|
| button = gr.Button("Genera") |
|
|
| output = gr.Textbox( |
| label="Risposta", |
| lines=12 |
| ) |
|
|
| button.click( |
| fn=generate, |
| inputs=[ |
| prompt, |
| max_tokens, |
| temperature, |
| top_p |
| ], |
| outputs=output |
| ) |
|
|
| demo.launch() |