dinghface's picture
Upload app.py with huggingface_hub
d40fc8d verified
Raw History Blame Contribute Delete
4.06 kB
import gradio as gr
from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM
import torch
MODEL_NAME = "dinghface/olmo3-190m-zh-full-continue"
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
torch_dtype=torch.bfloat16,
device_map="auto",
)
pipe = pipeline("text-generation", model=model, tokenizer=tokenizer)
def generate(
prompt,
max_new_tokens,
temperature,
top_p,
top_k,
repetition_penalty,
do_sample,
):
output = pipe(
prompt,
max_new_tokens=int(max_new_tokens),
do_sample=do_sample,
temperature=temperature if do_sample else None,
top_p=top_p if do_sample else None,
top_k=int(top_k) if do_sample else None,
repetition_penalty=repetition_penalty,
)
return output[0]["generated_text"]
EXAMPLES = [
["从前有座山,山里有座庙,", 256, 0.8, 0.9, 50, 1.2, True],
["人工智能是", 256, 0.7, 0.9, 50, 1.2, True],
["今天天气不错,我准备", 256, 0.8, 0.9, 50, 1.2, True],
["Python 是一种", 256, 0.7, 0.9, 50, 1.2, True],
["春天来了,万物复苏,", 256, 0.9, 0.95, 50, 1.1, True],
["在很久很久以前,", 256, 0.85, 0.9, 40, 1.2, True],
["The meaning of life is", 256, 0.8, 0.9, 50, 1.2, True],
["deep learning is", 256, 0.7, 0.9, 50, 1.2, True],
]
with gr.Blocks(title="OLMo3-190M-zh Continue Pretrain Demo") as demo:
gr.Markdown(
"""
# OLMo3-190M-zh 持续预训练 Demo
基于 [OLMo3-190M-zh-full](https://huggingface.co/dinghface/olmo3-190m-zh-full) 进行持续预训练的 190M 参数中文模型。
输入一段文字,模型会自动续写。
"""
)
with gr.Row():
with gr.Column(scale=2):
prompt = gr.Textbox(
label="输入提示词",
placeholder="在这里输入文字,模型会继续往下写...",
lines=5,
)
output = gr.Textbox(label="生成结果", lines=10)
with gr.Row():
submit_btn = gr.Button("生成", variant="primary")
clear_btn = gr.Button("清空")
with gr.Column(scale=1):
do_sample = gr.Checkbox(label="启用采样(关闭则为贪心解码)", value=True)
max_new_tokens = gr.Slider(
minimum=16, maximum=1024, value=256, step=16, label="最大生成长度"
)
temperature = gr.Slider(
minimum=0.1, maximum=2.0, value=0.8, step=0.05, label="Temperature(温度)"
)
top_p = gr.Slider(
minimum=0.1, maximum=1.0, value=0.9, step=0.05, label="Top-p(核采样)"
)
top_k = gr.Slider(
minimum=1, maximum=200, value=50, step=1, label="Top-k"
)
repetition_penalty = gr.Slider(
minimum=1.0, maximum=2.0, value=1.2, step=0.05, label="重复惩罚"
)
gr.Examples(
examples=EXAMPLES,
inputs=[prompt, max_new_tokens, temperature, top_p, top_k, repetition_penalty, do_sample],
)
gr.Markdown(
"""
### 参数说明
- **Temperature**:越高越随机,越低越确定。1.0 为默认,<1 更保守,>1 更有创意
- **Top-p**:核采样,从累积概率达到该值的 token 中采样。1.0 不过滤
- **Top-k**:只从概率最高的 k 个 token 中采样。值越大选择越多
- **重复惩罚**:>1 时惩罚重复内容,避免循环输出
- **启用采样**:关闭后使用贪心解码(每次选概率最高的 token),输出确定但单一
"""
)
submit_btn.click(
fn=generate,
inputs=[prompt, max_new_tokens, temperature, top_p, top_k, repetition_penalty, do_sample],
outputs=output,
)
clear_btn.click(fn=lambda: ("", ""), inputs=None, outputs=[prompt, output])
if __name__ == "__main__":
demo.launch()