Spaces:
Paused
Paused
| 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", | |
| ) | |
| 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}") | |