mps-models-1 / app.py
mps's picture
Upload 5 files
b2d2e14 verified
Raw
History Blame Contribute Delete
2.65 kB
import torch
import streamlit as st
from transformers import AutoModelForCausalLM, AutoTokenizer
MODEL_ID = "mps/blue-scrub-150M"
st.set_page_config(
page_title="Blue Scrub 150M",
page_icon="馃┖",
layout="wide",
)
@st.cache_resource(show_spinner="Loading model from Hugging Face Hub...")
def load_model():
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.float32,
low_cpu_mem_usage=True,
)
model.eval()
return tokenizer, model
def generate_text(prompt: str, max_new_tokens: int, temperature: float, top_p: float) -> str:
tokenizer, model = load_model()
inputs = tokenizer(prompt, return_tensors="pt")
do_sample = temperature > 0
generation_kwargs = {
"max_new_tokens": max_new_tokens,
"do_sample": do_sample,
"pad_token_id": tokenizer.eos_token_id,
"eos_token_id": tokenizer.eos_token_id,
}
if do_sample:
generation_kwargs.update({"temperature": temperature, "top_p": top_p})
with torch.no_grad():
output_ids = model.generate(**inputs, **generation_kwargs)
generated_ids = output_ids[0][inputs["input_ids"].shape[-1]:]
return tokenizer.decode(generated_ids, skip_special_tokens=True).strip()
st.title("馃┖ Blue Scrub 150M Inference")
st.caption(f"Model: `{MODEL_ID}` 路 Streamlit CPU Space 路 non-instruction-tuned base model")
st.warning(
"This is a base/non-instruction-tuned model. It completes text and is not a chat assistant. "
"Do not use outputs as medical advice. Free CPU inference can be slow."
)
with st.sidebar:
st.header("Generation settings")
max_new_tokens = st.slider("Max new tokens", min_value=8, max_value=256, value=96, step=8)
temperature = st.slider("Temperature", min_value=0.0, max_value=2.0, value=0.7, step=0.1)
top_p = st.slider("Top-p", min_value=0.05, max_value=1.0, value=0.9, step=0.05)
st.markdown("---")
st.markdown("[Open model card](https://huggingface.co/mps/blue-scrub-150M)")
prompt = st.text_area(
"Prompt",
value="Medical evidence suggests that",
height=180,
help="Use continuation-style prompts because this is a base model, not an instruction model.",
)
if st.button("Generate", type="primary", disabled=not prompt.strip()):
try:
with st.spinner("Generating..."):
text = generate_text(prompt, max_new_tokens, temperature, top_p)
st.subheader("Generated continuation")
st.write(text or "No text generated.")
except Exception as exc:
st.error(f"Inference failed: {exc}")