Spaces:
Paused
Paused
File size: 2,647 Bytes
b2d2e14 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | 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}")
|