ssdataanalysis's picture
Replace api_name=False with explicit private endpoints to avoid FnIndex errors
4d0d04c verified
Raw
History Blame Contribute Delete
57.5 kB
from __future__ import annotations
import json
import os
import queue
import threading
import time
from datetime import date
from pathlib import Path
from typing import Any
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("HF_DEACTIVATE_ASYNC_LOAD", "1")
try:
DEFAULT_COLD_START_DURATION_SECONDS = int(os.getenv("MUSE_COLD_START_DURATION_SECONDS", "120"))
except (TypeError, ValueError):
DEFAULT_COLD_START_DURATION_SECONDS = 120
SKIP_MODEL_LOAD = os.getenv("MUSE_SKIP_MODEL_LOAD", "0") == "1"
try:
import spaces
except ModuleNotFoundError:
if not SKIP_MODEL_LOAD:
raise
class _LocalSpaces:
@staticmethod
def GPU(*_args, **_kwargs):
def decorator(function):
return function
return decorator
spaces = _LocalSpaces()
import gradio as gr
from PIL import Image, ImageOps
import torch
from transformers import StoppingCriteria, StoppingCriteriaList, TextIteratorStreamer
from muse_core import (
APP_INPUT_TOKEN_LIMIT,
DEFAULT_MAX_NEW_TOKENS,
DEFAULT_REPETITION_PENALTY,
DEFAULT_SEED,
DEFAULT_TEMPERATURE,
DEFAULT_TOP_K,
DEFAULT_TOP_P,
MAX_NEW_TOKENS,
META_SAMPLING,
MODEL_CONTEXT_TOKENS,
NATIVE_GREEDY,
PRESETS,
choose_seed,
coerce_parsed_reply,
estimate_gpu_duration,
friendly_error,
generation_kwargs,
preset_values,
render_reply,
validate_controls,
)
MODEL_ID = "meta-models/Muse-Glimmer-30B"
MODEL_REVISION = "f84ecc3a0ea984a4c04542a84269e3d065350a6e"
ASSISTANT_MODEL_ID = "meta-models/Muse-Glimmer-30B-assistant"
ASSISTANT_MODEL_REVISION = "2c86316d689027b91123638739743fef1d425233"
EXPECTED_MODEL_TYPE = "muse_glimmer"
EXPECTED_MODEL_TYPES = {EXPECTED_MODEL_TYPE, "muse_glimmer_assistant"}
ASSISTANT_EXPECTED_PARAMETER_COUNT = 2_555_985_152
MODEL_CHOICES = [
("Muse Glimmer 30B (full BF16)", MODEL_ID),
("Muse Glimmer 30B-assistant (compact)", ASSISTANT_MODEL_ID),
]
SUBMIT_API_NAME = "chat_submit"
MODEL_NAME_ALIASES = {
"/Muse-Glimmer 30B": MODEL_ID,
"/Muse-Glimmer-30B": MODEL_ID,
"Muse Glimmer 30B (full BF16)": MODEL_ID,
"/Muse-Glimmer 30B-assistant": ASSISTANT_MODEL_ID,
"/Muse-Glimmer-30B-assistant": ASSISTANT_MODEL_ID,
"Muse Glimmer 30B-assistant (compact)": ASSISTANT_MODEL_ID,
"0": MODEL_ID,
"1": ASSISTANT_MODEL_ID,
"": MODEL_ID,
}
_BASE_MODEL_PATH = Path(os.getenv("MUSE_MODEL_PATH", "/models/muse-glimmer"))
_BASE_ASSISTANT_MODEL_PATH = Path(
os.getenv("MUSE_ASSISTANT_MODEL_PATH", "/models/muse-glimmer-assistant")
)
def _has_model_manifest(path: Path) -> bool:
has_config = (path / "config.json").is_file()
if not has_config:
return False
return (path / "chat_template.jinja").is_file() or (path / "tokenizer.json").is_file()
def _resolve_mount_path(model_root: Path) -> Path:
"""Handle both direct and mounted-directory layouts for model checkpoints."""
try:
if _has_model_manifest(model_root):
return model_root
except OSError:
return model_root
if not model_root.is_dir():
return model_root
# hf mount may expose the checkpoint under a subfolder of the declared mount path.
candidate_children = []
try:
for child in model_root.iterdir():
if child.is_dir() and _has_model_manifest(child):
candidate_children.append(child)
except OSError:
return model_root
if len(candidate_children) == 1:
return candidate_children[0]
if len(candidate_children) > 1:
# Prefer a child that explicitly references a Muse-Glimmer checkpoint name.
for child in candidate_children:
if child.name.startswith("Muse-Glimmer-30B"):
return child
return model_root
MODEL_REGISTRY = {
MODEL_ID: {
"revision": MODEL_REVISION,
"path": _resolve_mount_path(_BASE_MODEL_PATH),
"expected_model_type": "muse_glimmer",
"expected_parameter_count": 29_776_626_688,
"display": "Muse Glimmer 30B (full BF16)",
},
ASSISTANT_MODEL_ID: {
"revision": ASSISTANT_MODEL_REVISION,
"path": _resolve_mount_path(_BASE_ASSISTANT_MODEL_PATH),
"expected_model_type": "muse_glimmer_assistant",
"expected_parameter_count": ASSISTANT_EXPECTED_PARAMETER_COUNT,
"display": "Muse Glimmer 30B-assistant (compact)",
},
}
def _coerce_model_id(model_id: Any) -> str:
if model_id in (None, []):
return MODEL_DEFAULT_ID
if isinstance(model_id, (tuple, list)):
if not model_id:
return MODEL_DEFAULT_ID
if len(model_id) > 1 and isinstance(model_id[1], str):
return model_id[1]
if isinstance(model_id[0], str):
return _coerce_model_id(model_id[0])
return MODEL_DEFAULT_ID
if isinstance(model_id, int):
choices = [value for _label, value in MODEL_CHOICES]
if 0 <= model_id < len(choices):
return choices[model_id]
return MODEL_DEFAULT_ID
if isinstance(model_id, str):
normalized = model_id.strip()
if normalized in MODEL_NAME_ALIASES:
return MODEL_NAME_ALIASES[normalized]
if model_id.isdigit():
choices = [value for _label, value in MODEL_CHOICES]
idx = int(model_id)
if 0 <= idx < len(choices):
return choices[idx]
return model_id
return str(model_id)
def _resolve_default_model_id() -> str:
configured = os.getenv("MUSE_DEFAULT_MODEL_ID", MODEL_ID)
if configured not in MODEL_REGISTRY:
configured = MODEL_ID
configured_path = MODEL_REGISTRY[configured]["path"]
if configured_path.is_dir():
return configured
for model_id, spec in MODEL_REGISTRY.items():
if model_id == configured:
continue
if spec["path"].is_dir():
return model_id
return configured
MODEL_DEFAULT_ID = _resolve_default_model_id()
MAX_HISTORY_MESSAGES = 20
MAX_HISTORY_IMAGES = 2
MAX_IMAGE_EDGE = 2_048
MAX_IMAGE_PIXELS = 4_194_304
ACTIVE_MODEL_ID: str | None = None
ACTIVE_MODEL = None
ACTIVE_PROCESSOR = None
def _model_spec(model_id: str) -> dict[str, Any]:
if model_id not in MODEL_REGISTRY:
raise ValueError(f"Unknown model selection: {model_id}")
return MODEL_REGISTRY[model_id]
def _is_model_checkpoint(path: str | os.PathLike[str], model_path: Path) -> bool:
try:
candidate = Path(path).resolve()
model_root = model_path.resolve()
except (OSError, RuntimeError, ValueError):
return False
candidate_text = str(candidate)
model_root_text = str(model_root)
return candidate == model_root or candidate_text.startswith(model_root_text + os.sep)
def _normalize_load_result(result: Any) -> tuple[Any, dict[str, Any]]:
if isinstance(result, tuple):
if len(result) >= 2:
return result[0], result[1]
return result[0], {}
if isinstance(result, dict):
return result.get("model"), result
return result, {}
def _supports_generation(model: Any) -> bool:
return callable(getattr(model, "generate", None))
def _load_model_with_pread(
model_class,
model_path: Path,
*,
use_safetensors: bool = True,
safe_open_backend: str | None = "pread",
trust_remote_code: bool = False,
):
"""Load the mounted shards sequentially without mmap or whole-shard RAM copies.
Transformers 5.15 deliberately disables mmap for Hugging Face model volumes because
concurrent page faults can deadlock hf-mount. Its fallback reads an entire safetensors
shard into host RAM; Muse Glimmer's first shard is about 50 GB, while a standard Space
has far less host RAM. Safetensors 0.8's pread backend avoids both failure modes and lets
Transformers materialize and dispatch one tensor at a time.
"""
from safetensors import safe_open as safetensors_safe_open
from transformers import modeling_utils
if not hasattr(modeling_utils, "_is_on_hf_mount") or not hasattr(modeling_utils, "safe_open"):
return _load_model_direct(model_class, model_path)
shards = sorted(model_path.glob("*.safetensors"))
if not shards:
raise RuntimeError("The mounted checkpoint contains no safetensors shards.")
# Fail early with a small header-only read for the selected backend.
safe_open_kwargs = {"framework": "pt", "device": "cpu"}
if safe_open_backend is not None:
safe_open_kwargs["backend"] = safe_open_backend
with safetensors_safe_open(str(shards[0]), **safe_open_kwargs) as checkpoint:
first_key = next(iter(checkpoint.keys()), None)
if first_key is None:
raise RuntimeError("The mounted safetensors checkpoint is empty.")
checkpoint.get_slice(first_key).get_shape()
original_mount_check = modeling_utils._is_on_hf_mount
original_safe_open = modeling_utils.safe_open
def model_mount_check(path):
if _is_model_checkpoint(path, model_path):
return False
return original_mount_check(path)
def model_safe_open(path, *args, **kwargs):
if _is_model_checkpoint(path, model_path) and os.fspath(path).endswith(".safetensors"):
if safe_open_backend is not None:
kwargs["backend"] = safe_open_backend
return original_safe_open(path, *args, **kwargs)
modeling_utils._is_on_hf_mount = model_mount_check
modeling_utils.safe_open = model_safe_open
try:
loaded = model_class.from_pretrained(
model_path,
dtype=torch.bfloat16,
device_map={"": "cuda"},
local_files_only=True,
trust_remote_code=trust_remote_code,
attn_implementation="sdpa",
output_loading_info=True,
disable_mmap=False,
use_safetensors=use_safetensors,
)
return _normalize_load_result(loaded)
finally:
modeling_utils._is_on_hf_mount = original_mount_check
modeling_utils.safe_open = original_safe_open
def _load_model_direct(
model_class,
model_path: Path,
*,
use_safetensors: bool = True,
trust_remote_code: bool = False,
):
return _normalize_load_result(
model_class.from_pretrained(
model_path,
dtype=torch.bfloat16,
device_map={"": "cuda"},
local_files_only=True,
trust_remote_code=trust_remote_code,
attn_implementation="sdpa",
output_loading_info=True,
use_safetensors=use_safetensors,
)
)
def _load_model_candidate(
model_class,
model_path: Path,
*,
trust_remote_code: bool,
):
for use_safetensors in (True, False):
for safe_open_backend in ("pread", "read", None):
try:
return _load_model_with_pread(
model_class,
model_path,
use_safetensors=use_safetensors,
safe_open_backend=safe_open_backend,
trust_remote_code=trust_remote_code,
)
except Exception:
pass
return _load_model_direct(model_class, model_path, use_safetensors=False, trust_remote_code=trust_remote_code)
def _load_model_candidate_or_remote(
model_class,
spec: dict[str, Any],
model_id: str,
*,
trust_remote_code: bool,
):
model_path = spec["path"]
revision = spec["revision"]
try:
return _load_model_candidate(
model_class,
model_path,
trust_remote_code=trust_remote_code,
)
except Exception:
pass
for use_safetensors in (True, False):
try:
return _normalize_load_result(
model_class.from_pretrained(
model_id,
revision=revision,
dtype=torch.bfloat16,
device_map={"": "cuda"},
local_files_only=False,
trust_remote_code=trust_remote_code,
attn_implementation="sdpa",
output_loading_info=True,
use_safetensors=use_safetensors,
cache_dir="/tmp/huggingface-model-cache",
)
)
except Exception:
pass
raise RuntimeError("Unable to load the selected checkpoint from local mount or remote Hub download.")
def _load_runtime(model_id: str):
spec = _model_spec(model_id)
model_path = spec["path"]
revision = spec["revision"]
expected_model_type = spec["expected_model_type"]
has_mount = model_path.is_dir()
has_assistant_fallback_mount = MODEL_REGISTRY[MODEL_ID]["path"].is_dir()
use_remote = not has_mount and model_id == ASSISTANT_MODEL_ID
if not has_mount and not use_remote:
raise RuntimeError(
f"The selected Muse Glimmer full model mount is missing at {model_path}. "
"Attach the read-only model volume before starting the Space."
)
from transformers import AutoConfig, AutoModelForCausalLM, AutoProcessor, AutoTokenizer
print(
f"[startup] Loading processor from "
f"{'model repository' if use_remote else model_path} ({revision[:12]}…).",
flush=True,
)
source = model_id if use_remote else model_path
processor_kwargs = {
"revision": revision,
"local_files_only": not use_remote,
"trust_remote_code": False,
}
config_kwargs = {
"revision": revision,
"local_files_only": not use_remote,
"trust_remote_code": False,
}
if model_id == ASSISTANT_MODEL_ID:
processor_source = (
MODEL_REGISTRY[MODEL_ID]["path"] if has_assistant_fallback_mount else MODEL_ID
)
if processor_source == MODEL_REGISTRY[MODEL_ID]["path"]:
processor_kwargs["local_files_only"] = True
processor_kwargs["revision"] = MODEL_REVISION
else:
processor_kwargs["local_files_only"] = False
processor_kwargs["revision"] = MODEL_REVISION
print(
"[startup] Assistant-selected checkpoint will reuse "
f"base-tokenization assets from `{processor_source}`.",
flush=True,
)
else:
processor_source = source
if use_remote:
processor_kwargs["cache_dir"] = "/tmp/huggingface-model-cache"
config_kwargs["cache_dir"] = "/tmp/huggingface-model-cache"
if model_id == ASSISTANT_MODEL_ID:
try:
processor = AutoProcessor.from_pretrained(processor_source, **processor_kwargs)
tokenizer = AutoTokenizer.from_pretrained(processor_source, **processor_kwargs)
except Exception:
print(
"[startup] Processor loading failed without trust_remote_code; retrying with trust_remote_code=True.",
flush=True,
)
fallback_processor_kwargs = dict(processor_kwargs)
fallback_processor_kwargs["trust_remote_code"] = True
processor = AutoProcessor.from_pretrained(processor_source, **fallback_processor_kwargs)
tokenizer = AutoTokenizer.from_pretrained(processor_source, **fallback_processor_kwargs)
if not hasattr(processor, "tokenizer"):
processor.tokenizer = tokenizer
else:
try:
processor = AutoProcessor.from_pretrained(source, **processor_kwargs)
except Exception:
print(
"[startup] Processor loading failed without trust_remote_code; retrying with trust_remote_code=True.",
flush=True,
)
fallback_processor_kwargs = dict(processor_kwargs)
fallback_processor_kwargs["trust_remote_code"] = True
processor = AutoProcessor.from_pretrained(source, **fallback_processor_kwargs)
try:
config = AutoConfig.from_pretrained(source, **config_kwargs)
except Exception:
config = None
if config is None:
fallback_config_kwargs = dict(config_kwargs)
fallback_config_kwargs["trust_remote_code"] = True
try:
config = AutoConfig.from_pretrained(source, **fallback_config_kwargs)
print(
"[startup] AutoConfig with trust_remote_code succeeded for the selected checkpoint.",
flush=True,
)
except Exception as error:
raise RuntimeError("Unable to load model configuration from the selected checkpoint.") from error
model_type = getattr(config, "model_type", None)
if model_type != expected_model_type:
print(
f"[startup] Warning: checkpoint model_type={model_type} while expected {expected_model_type}. "
"Proceeding with detected architecture checks.",
flush=True,
)
if model_type == "muse_glimmer":
from transformers import MuseGlimmerForConditionalGeneration
model_candidates = ((MuseGlimmerForConditionalGeneration, False),)
elif model_type == "muse_glimmer_assistant":
try:
from transformers.models.muse_glimmer_assistant.modeling_muse_glimmer_assistant import (
MuseGlimmerAssistantModel,
)
model_candidates = (
(MuseGlimmerAssistantModel, False),
(MuseGlimmerAssistantModel, True),
)
except Exception:
model_candidates = (
(AutoModelForCausalLM, False),
(AutoModelForCausalLM, True),
)
else:
raise RuntimeError(
f"Unsupported model type from checkpoint: {model_type}. "
f"Expected {expected_model_type or 'a Muse Glimmer variant'}."
)
print("[startup] Loading the selected Muse Glimmer checkpoint onto ZeroGPU.", flush=True)
loading_info = {}
loading_error = None
model = None
used_model_class = None
try:
for model_class, trust_remote_code in model_candidates:
used_model_class = getattr(model_class, "__name__", str(model_class))
try:
if model_type == "muse_glimmer_assistant":
print(
f"[startup] Trying {used_model_class} for assistant checkpoint "
f"with trust_remote_code={trust_remote_code}.",
flush=True,
)
model, loading_info = _load_model_candidate_or_remote(
model_class,
spec,
model_id,
trust_remote_code=trust_remote_code,
)
loading_error = None
break
except Exception as error:
loading_error = error
print(
f"[startup] {used_model_class} load failed ({type(error).__name__}); trying next option if available.",
flush=True,
)
if model is None:
raise RuntimeError(f"No compatible loader could initialize model class for `{model_id}`.")
except Exception as error: # pragma: no cover - runtime-only edge
if loading_error is None:
loading_error = error
raise
if not isinstance(loading_info, dict):
loading_info = {}
loading_failures = {
key: loading_info.get(key)
for key in (
"missing_keys",
"unexpected_keys",
"mismatched_keys",
"conversion_errors",
"error_msgs",
)
if loading_info.get(key)
}
if loading_failures:
raise RuntimeError(
"The pinned checkpoint did not load cleanly: "
+ ", ".join(f"{key}={len(value)}" for key, value in loading_failures.items())
)
if model_type == "muse_glimmer_assistant" and not _supports_generation(model):
print(
"[startup] Loaded assistant checkpoint is not a standalone generator; inference will fallback "
"to the full model at request time when selected.",
flush=True,
)
loaded_model_type = getattr(model.config, "model_type", None)
if loaded_model_type is not None and loaded_model_type not in EXPECTED_MODEL_TYPES:
raise RuntimeError("The selected checkpoint is not a Muse Glimmer model.")
if loaded_model_type is None:
print("[startup] Checkpoint config has no model_type; proceeding with expected loader class.", flush=True)
parameter_count = sum(parameter.numel() for parameter in model.parameters())
expected_parameter_count = spec["expected_parameter_count"]
if expected_parameter_count is not None and parameter_count != expected_parameter_count:
raise RuntimeError(
f"Unexpected parameter count: {parameter_count:,}; expected {expected_parameter_count:,}."
)
model.eval()
if loading_error is not None:
print(f"[startup] Loaded with fallback loader after: {type(loading_error).__name__}", flush=True)
print(
f"[startup] Ready: {parameter_count:,} parameters from `{model_id}` ({revision[:12]}…).",
flush=True,
)
return processor, model
def _activate_model(model_id: str):
global ACTIVE_MODEL_ID, ACTIVE_MODEL, ACTIVE_PROCESSOR, PROCESSOR, MODEL
if model_id not in MODEL_REGISTRY:
raise ValueError(f"Unknown model selection: {model_id}")
if ACTIVE_MODEL_ID == model_id and ACTIVE_MODEL is not None and ACTIVE_PROCESSOR is not None:
return ACTIVE_MODEL, ACTIVE_PROCESSOR
if ACTIVE_MODEL is not None:
del ACTIVE_MODEL
if ACTIVE_PROCESSOR is not None:
del ACTIVE_PROCESSOR
if torch.cuda.is_available():
torch.cuda.empty_cache()
ACTIVE_PROCESSOR, ACTIVE_MODEL = _load_runtime(model_id)
ACTIVE_MODEL_ID = model_id
if torch.cuda.is_available():
torch.cuda.synchronize()
PROCESSOR = ACTIVE_PROCESSOR
MODEL = ACTIVE_MODEL
return ACTIVE_PROCESSOR, ACTIVE_MODEL
if SKIP_MODEL_LOAD:
PROCESSOR = None
MODEL = None
else:
available_models = [model_id for model_id, spec in MODEL_REGISTRY.items() if spec["path"].is_dir()]
if available_models:
print(
f"[startup] Model loading deferred until first request. Available mounts: {', '.join(available_models)}",
flush=True,
)
else:
print("[startup] No checkpoint mounts are available at startup; model loading is deferred.", flush=True)
PROCESSOR = None
MODEL = None
class _StopOnEvent(StoppingCriteria):
def __init__(self, event: threading.Event):
self.event = event
def __call__(self, input_ids, scores, **kwargs):
del scores, kwargs
return torch.full(
(input_ids.shape[0],),
self.event.is_set(),
dtype=torch.bool,
device=input_ids.device,
)
def _coerce_image_input(image: Any) -> Image.Image | None:
if image is None or (isinstance(image, str) and not image):
return None
if not isinstance(image, Image.Image):
raise ValueError("The image upload could not be decoded.")
return _normalize_image(image)
def _normalize_image(image: Image.Image | None) -> Image.Image | None:
width, height = image.size
if width < 1 or height < 1:
raise ValueError("The image has invalid dimensions.")
if width * height > MAX_IMAGE_PIXELS:
scale = (MAX_IMAGE_PIXELS / float(width * height)) ** 0.5
image = image.resize(
(max(1, int(width * scale)), max(1, int(height * scale))),
Image.Resampling.LANCZOS,
)
image = ImageOps.exif_transpose(image)
image.thumbnail((MAX_IMAGE_EDGE, MAX_IMAGE_EDGE), Image.Resampling.LANCZOS)
clean = Image.new("RGB", image.size)
if image.mode == "RGBA":
background = Image.new("RGBA", image.size, "white")
background.alpha_composite(image)
clean.paste(background.convert("RGB"))
else:
clean.paste(image.convert("RGB"))
return clean
def _response_tokenizer_for(obj: Any):
tokenizer = getattr(obj, "tokenizer", None)
if tokenizer is not None:
return tokenizer
return getattr(obj, "_tokenizer", None)
def _coerce_chat_objects(processor_or_tokenizer: Any, model_id: str) -> tuple[Any, Any]:
"""Return a processor/tokenizer pair that both support templating and parser wiring.
This guards against edge cases where processor loading returns an unexpected object
(for example during Transformers internals or runtime cache fallback behavior).
"""
from transformers import AutoProcessor, AutoTokenizer
spec = _model_spec(model_id)
source = spec["path"] if spec["path"].is_dir() else model_id
base_kwargs = {
"revision": spec["revision"],
"local_files_only": source == spec["path"] and spec["path"].is_dir(),
"trust_remote_code": False,
}
candidates: list[Any] = [processor_or_tokenizer]
tokenized = _response_tokenizer_for(processor_or_tokenizer)
if tokenized is not None:
candidates.append(tokenized)
def _supports_template(candidate: Any) -> bool:
return candidate is not None and hasattr(candidate, "apply_chat_template")
def _valid(candidate: Any) -> bool:
return _supports_template(candidate) and hasattr(candidate, "get_response_parser")
for candidate in candidates:
if candidate is not None and _valid(candidate):
return candidate, _response_tokenizer_for(candidate) or candidate
for candidate in candidates:
if _supports_template(candidate):
return candidate, _response_tokenizer_for(candidate) or candidate
for trust_remote_code in (False, True):
fallback_kwargs = dict(base_kwargs)
fallback_kwargs["trust_remote_code"] = trust_remote_code
try:
candidate = AutoProcessor.from_pretrained(source, **fallback_kwargs)
if _valid(candidate):
return candidate, _response_tokenizer_for(candidate) or candidate
except Exception:
pass
try:
candidate = AutoTokenizer.from_pretrained(source, **fallback_kwargs)
if _valid(candidate):
return candidate, candidate
if _supports_template(candidate):
return candidate, candidate
except Exception:
pass
raise RuntimeError("Unable to initialize chat template/parser components for the selected model.")
def _parse_llm_response(text: str | None) -> tuple[str, str]:
text = (text or "").strip()
if not text:
return "", ""
think_open = "<think>"
think_close = "</think>"
start = text.find(think_open)
if start == -1:
return "", text
start += len(think_open)
close = text.find(think_close, start)
if close == -1:
return text[start:].strip(), ""
reasoning = text[start:close].strip()
content = text[close + len(think_close) :].strip()
return reasoning, content
def _user_content(prompt: str, image: Image.Image | None):
if image is None:
return prompt
return [
{"type": "image", "image": image},
{"type": "text", "text": prompt},
]
def _visible_user_message(prompt: str, image: Image.Image | None) -> str:
if image is None:
return prompt
return f"{prompt}\n\n_🖼️ Image attached to this turn._"
def _clean_model_history(history) -> list[dict[str, Any]]:
cleaned: list[dict[str, Any]] = []
for message in list(history or [])[-MAX_HISTORY_MESSAGES:]:
if not isinstance(message, dict) or message.get("role") not in {"user", "assistant"}:
continue
if "content" not in message:
continue
safe = {"role": message["role"], "content": message["content"]}
if message["role"] == "assistant" and isinstance(message.get("reasoning_content"), str):
safe["reasoning_content"] = message["reasoning_content"]
cleaned.append(safe)
if cleaned and cleaned[0]["role"] == "assistant":
cleaned.pop(0)
# Preserve recent multimodal context without repeatedly serializing an unbounded
# number of raw PIL objects through Gradio State/ZeroGPU IPC.
kept_images = 0
for message in reversed(cleaned):
content = message.get("content")
if message.get("role") != "user" or not isinstance(content, list):
continue
has_image = any(isinstance(part, dict) and part.get("type") == "image" for part in content)
if not has_image:
continue
kept_images += 1
if kept_images <= MAX_HISTORY_IMAGES:
continue
text_parts = [
part.get("text", "")
for part in content
if isinstance(part, dict) and part.get("type") == "text"
]
message["content"] = "\n".join(part for part in text_parts if part).strip()
return cleaned
def _apply_template(processor, messages: list[dict[str, Any]], reasoning_strength: str):
return processor.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
reasoning_strength=reasoning_strength,
current_date=date.today().isoformat(),
return_dict=True,
return_tensors="pt",
)
def _prepare_inputs(
processor,
model_history,
prompt: str,
image: Image.Image | None,
system_prompt: str,
reasoning_strength: str,
max_new_tokens: int,
):
retained = _clean_model_history(model_history)
current_user = {"role": "user", "content": _user_content(prompt, image)}
trimmed_messages = 0
while True:
messages: list[dict[str, Any]] = []
if system_prompt.strip():
messages.append({"role": "system", "content": system_prompt.strip()})
messages.extend(retained)
messages.append(current_user)
encoded = _apply_template(processor, messages, reasoning_strength)
input_tokens = int(encoded["input_ids"].shape[-1])
if input_tokens <= APP_INPUT_TOKEN_LIMIT:
break
if not retained:
raise ValueError(
f"The current turn exceeds the app input limit of {APP_INPUT_TOKEN_LIMIT:,} tokens."
)
retained.pop(0)
trimmed_messages += 1
if retained and retained[0].get("role") == "assistant":
retained.pop(0)
trimmed_messages += 1
if input_tokens + int(max_new_tokens) > MODEL_CONTEXT_TOKENS:
raise ValueError("The prompt and response budget exceed the model context window.")
return retained, current_user, encoded, input_tokens, trimmed_messages
def _move_inputs_to_model(model, encoded):
device = next(model.parameters()).device
moved = {}
for key, value in encoded.items():
if not torch.is_tensor(value):
moved[key] = value
continue
value = value.to(device)
if key in {"pixel_values", "pixel_values_videos"} and value.is_floating_point():
value = value.to(dtype=torch.bfloat16)
moved[key] = value
return moved
def _gpu_duration(
prompt,
image,
selected_model,
chat_history,
model_history,
system_prompt,
reasoning_strength,
do_sample,
max_new_tokens,
temperature,
top_p,
top_k,
repetition_penalty,
seed,
randomize_seed,
show_reasoning,
):
_ = (
prompt,
chat_history,
system_prompt,
do_sample,
temperature,
top_p,
top_k,
repetition_penalty,
seed,
randomize_seed,
show_reasoning,
)
selected_model = _coerce_model_id(selected_model) or MODEL_DEFAULT_ID
try:
max_new_tokens = int(max_new_tokens)
except Exception:
max_new_tokens = DEFAULT_MAX_NEW_TOKENS
has_image = False
try:
has_image = _coerce_image_input(image) is not None
except ValueError:
has_image = False
needs_warmup = selected_model != ACTIVE_MODEL_ID
estimated = estimate_gpu_duration(max_new_tokens, has_image)
if selected_model == MODEL_ID and needs_warmup:
return min(estimated, DEFAULT_COLD_START_DURATION_SECONDS)
if selected_model == ASSISTANT_MODEL_ID:
if ACTIVE_MODEL_ID == MODEL_ID:
return estimated
return min(estimated, DEFAULT_COLD_START_DURATION_SECONDS)
return estimated
def _format_status(
*,
phase: str,
selected_model: str,
input_tokens: int,
output_tokens: int,
elapsed: float,
used_seed: int,
do_sample: bool,
trimmed_messages: int,
) -> str:
mode = "sampling" if do_sample else "native greedy"
trimmed = f" · trimmed {trimmed_messages} old messages" if trimmed_messages else ""
return (
f"{phase} · {input_tokens:,} input / {output_tokens:,} output tokens · "
f"{elapsed:.1f}s · {mode} · seed {used_seed}{trimmed} · {selected_model}"
)
@spaces.GPU(size="xlarge", duration=_gpu_duration)
def _generate_turn(
prompt,
image,
selected_model,
chat_history,
model_history,
system_prompt,
reasoning_strength,
do_sample,
max_new_tokens,
temperature,
top_p,
top_k,
repetition_penalty,
seed,
randomize_seed,
show_reasoning,
):
original_chat = list(chat_history or [])
original_model_history = list(model_history or [])
generation_thread: threading.Thread | None = None
stop_event = threading.Event()
try:
selected_model = _coerce_model_id(selected_model) or MODEL_DEFAULT_ID
selected_model_name = _model_spec(selected_model).get("display", selected_model)
active_inference_model = selected_model
model_fallback = False
processor, model = _activate_model(selected_model)
if active_inference_model == ASSISTANT_MODEL_ID and not _supports_generation(model):
print(
"[inference] Assistant checkpoint does not expose generate(); falling back to full model for this request.",
flush=True,
)
model_fallback = True
active_inference_model = MODEL_ID
processor, model = _activate_model(active_inference_model)
selected_model_name = _model_spec(MODEL_ID).get("display", MODEL_ID)
if processor is None or model is None:
raise RuntimeError("Model loading is unavailable for this request.")
prompt = (prompt or "").strip()
if not prompt:
raise ValueError("Write a prompt before generating.")
if len(prompt) > 20_000:
raise ValueError("The prompt is too long; keep it below 20,000 characters.")
max_new_tokens = int(DEFAULT_MAX_NEW_TOKENS if max_new_tokens is None else max_new_tokens)
temperature = DEFAULT_TEMPERATURE if temperature is None else float(temperature)
top_p = DEFAULT_TOP_P if top_p is None else float(top_p)
top_k = DEFAULT_TOP_K if top_k is None else int(top_k)
repetition_penalty = (
DEFAULT_REPETITION_PENALTY
if repetition_penalty is None
else float(repetition_penalty)
)
validate_controls(
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
repetition_penalty=repetition_penalty,
reasoning_strength=reasoning_strength,
)
used_seed = choose_seed(seed, bool(randomize_seed))
clean_image = _coerce_image_input(image)
processor, response_parser_tokenizer = _coerce_chat_objects(
processor, active_inference_model
)
retained, current_user, encoded, input_tokens, trimmed_messages = _prepare_inputs(
processor,
original_model_history,
prompt,
clean_image,
system_prompt or "",
reasoning_strength,
max_new_tokens,
)
model_inputs = _move_inputs_to_model(model, encoded)
input_length = int(model_inputs["input_ids"].shape[-1])
prefix_ids = encoded["input_ids"][0].detach().cpu()
torch.manual_seed(used_seed)
torch.cuda.manual_seed_all(used_seed)
streamer = TextIteratorStreamer(
response_parser_tokenizer,
skip_prompt=True,
skip_special_tokens=False,
timeout=5.0,
)
parser = (
response_parser_tokenizer.get_response_parser(prefix=prefix_ids)
if hasattr(response_parser_tokenizer, "get_response_parser")
else None
)
buffers = {"reasoning_content": "", "content": ""}
streamed_chunks: list[str] = []
if parser is not None:
for event in parser.initial_events:
if event.get("type") == "region_chunk" and event.get("field") in buffers:
buffers[event["field"]] += event.get("text", "")
kwargs = {
**model_inputs,
**generation_kwargs(
do_sample=bool(do_sample),
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
repetition_penalty=repetition_penalty,
),
"streamer": streamer,
"stopping_criteria": StoppingCriteriaList([_StopOnEvent(stop_event)]),
"max_time": float(max(30, estimate_gpu_duration(max_new_tokens, clean_image is not None))),
}
errors: list[BaseException] = []
result_box: list[Any] = []
def run_model() -> None:
try:
with torch.inference_mode():
result_box.append(model.generate(**kwargs))
except BaseException as error:
errors.append(error)
streamer.on_finalized_text("", stream_end=True)
generation_thread = threading.Thread(target=run_model, daemon=True)
started = time.perf_counter()
generation_thread.start()
user_message = {"role": "user", "content": _visible_user_message(prompt, clean_image)}
working_chat = original_chat + [user_message]
last_yield = 0.0
while True:
try:
chunk = next(streamer)
except queue.Empty:
if not generation_thread.is_alive():
if errors:
break
raise RuntimeError("The generation stream ended unexpectedly.")
now = time.perf_counter()
yield (
working_chat
+ [
{
"role": "assistant",
"content": render_reply(
buffers["reasoning_content"],
buffers["content"],
show_reasoning=bool(show_reasoning),
pending=True,
),
}
],
gr.skip(),
gr.skip(),
gr.skip(),
gr.skip(),
_format_status(
selected_model=(
f"{selected_model_name} (assistant checkpoint fallback to full model)"
if model_fallback
else selected_model_name
),
phase="Generating",
input_tokens=input_tokens,
output_tokens=0,
elapsed=now - started,
used_seed=used_seed,
do_sample=bool(do_sample),
trimmed_messages=trimmed_messages,
),
)
last_yield = now
continue
except StopIteration:
break
if parser is not None:
for event in parser.feed(chunk):
field = event.get("field")
if field not in buffers:
continue
if event.get("type") == "region_chunk":
buffers[field] += event.get("text", "")
elif event.get("type") == "region_close" and isinstance(event.get("value"), str):
buffers[field] = event["value"]
else:
streamed_chunks.append(chunk)
reasoning, content = _parse_llm_response("".join(streamed_chunks))
buffers["reasoning_content"] = reasoning
buffers["content"] = content
now = time.perf_counter()
if now - last_yield < 0.06:
continue
partial = render_reply(
buffers["reasoning_content"],
buffers["content"],
show_reasoning=bool(show_reasoning),
pending=True,
)
elapsed = now - started
yield (
working_chat + [{"role": "assistant", "content": partial}],
gr.skip(),
gr.skip(),
gr.skip(),
gr.skip(),
_format_status(
selected_model=(
f"{selected_model_name} (assistant checkpoint fallback to full model)"
if model_fallback
else selected_model_name
),
phase="Generating",
input_tokens=input_tokens,
output_tokens=0,
elapsed=elapsed,
used_seed=used_seed,
do_sample=bool(do_sample),
trimmed_messages=trimmed_messages,
),
)
last_yield = now
generation_thread.join(timeout=3)
if generation_thread.is_alive():
raise RuntimeError(
"Generation exceeded its timeout envelope. "
"Lower the response budget and try again."
)
if errors:
raise errors[0]
if parser is not None:
parsed_message, final_events = parser.finalize()
for event in final_events:
field = event.get("field")
if (
field in buffers
and event.get("type") == "region_close"
and isinstance(event.get("value"), str)
):
buffers[field] = event["value"]
parsed = coerce_parsed_reply(parsed_message)
reasoning = parsed.reasoning or buffers["reasoning_content"].strip()
content = parsed.content or buffers["content"].strip()
else:
reasoning, content = _parse_llm_response("".join(streamed_chunks))
output_tokens = 0
ended_with_limit = False
if result_box:
generated = result_box[0]
output_tokens = int(generated.shape[-1]) - input_length
ended_with_limit = output_tokens >= max_new_tokens
if not reasoning and not content:
raise RuntimeError("The model returned no visible response fields.")
visible_reply = render_reply(
reasoning,
content,
show_reasoning=bool(show_reasoning),
hit_token_limit=ended_with_limit,
)
assistant_state = {"role": "assistant", "content": content}
if reasoning:
assistant_state["reasoning_content"] = reasoning
updated_model_history = retained + [current_user, assistant_state]
updated_chat = working_chat + [{"role": "assistant", "content": visible_reply}]
elapsed = time.perf_counter() - started
yield (
updated_chat,
updated_model_history,
updated_chat,
"",
None,
_format_status(
selected_model=(
f"{selected_model_name} (assistant checkpoint fallback to full model)"
if model_fallback
else selected_model_name
),
phase="Complete",
input_tokens=input_tokens,
output_tokens=output_tokens,
elapsed=elapsed,
used_seed=used_seed,
do_sample=bool(do_sample),
trimmed_messages=trimmed_messages,
),
)
except GeneratorExit:
raise
except BaseException as error:
print(f"[inference] {type(error).__name__}: {error}", flush=True)
if torch.cuda.is_available():
torch.cuda.empty_cache()
yield (
original_chat,
gr.skip(),
original_chat,
gr.skip(),
gr.skip(),
f"Error · {friendly_error(error)}",
)
finally:
stop_event.set()
if generation_thread is not None and generation_thread.is_alive():
generation_thread.join(timeout=3)
if generation_thread.is_alive():
print("[inference] Generation worker did not stop within grace window.", flush=True)
def _validate_generation_request(
prompt,
image,
selected_model,
chat_history,
model_history,
system_prompt,
reasoning_strength,
do_sample,
max_new_tokens,
temperature,
top_p,
top_k,
repetition_penalty,
seed,
randomize_seed,
show_reasoning,
):
del chat_history, model_history, do_sample, show_reasoning
valid = True
message = ""
try:
selected_model = _coerce_model_id(selected_model) or MODEL_DEFAULT_ID
spec = _model_spec(selected_model)
max_new_tokens = 32 if max_new_tokens is None else int(max_new_tokens)
temperature = 1.0 if temperature is None else float(temperature)
top_p = 0.95 if top_p is None else float(top_p)
top_k = 64 if top_k is None else int(top_k)
repetition_penalty = 1.0 if repetition_penalty is None else float(repetition_penalty)
image = _coerce_image_input(image)
if not spec["path"].is_dir() and selected_model != ASSISTANT_MODEL_ID:
raise ValueError(f"The selected model checkpoint is not mounted at {spec['path']}.")
if selected_model == ASSISTANT_MODEL_ID and not MODEL_REGISTRY[MODEL_ID]["path"].is_dir():
raise ValueError(
"Assistant checkpoint inference currently falls back to the full model, "
f"but the full model mount is missing at {MODEL_REGISTRY[MODEL_ID]['path']}."
)
prompt = (prompt or "").strip()
if not prompt:
raise ValueError("Write a prompt before generating.")
if len(prompt) > 20_000:
raise ValueError("The prompt is too long; keep it below 20,000 characters.")
if len(system_prompt or "") > 20_000:
raise ValueError("The system instruction is too long; keep it below 20,000 characters.")
validate_controls(
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
repetition_penalty=repetition_penalty,
reasoning_strength=reasoning_strength,
)
if not bool(randomize_seed):
choose_seed(seed, False)
except (TypeError, ValueError) as error:
valid = False
message = str(error)
verdicts = [gr.validate(valid, message)]
verdicts.extend(gr.validate(True, "") for _ in range(15))
return tuple(verdicts)
def _stop_conversation(chat_snapshot):
return list(chat_snapshot or []), "Stopped · the unfinished turn was not added to model history"
def _clear_conversation():
return [], [], [], "", None, None, "Ready · native greedy · reasoning high"
def _set_preset(name: str):
return preset_values(name)
CSS = """
:root {
--ink: #161225;
--muted: #686177;
--line: #e8e1f1;
--paper: #ffffff;
--wash: #faf8fd;
--violet: #6d28d9;
--cyan: #0e7490;
}
.gradio-container {
max-width: 1180px !important;
margin: 0 auto !important;
background:
radial-gradient(circle at 8% 0%, rgba(109, 40, 217, .12), transparent 31rem),
radial-gradient(circle at 92% 0%, rgba(14, 116, 144, .10), transparent 29rem),
var(--wash);
}
#hero {
padding: 26px 28px 22px;
border: 1px solid var(--line);
border-radius: 22px;
background: rgba(255, 255, 255, .90);
box-shadow: 0 18px 50px rgba(41, 24, 72, .07);
}
#hero h1 { margin-bottom: 7px; letter-spacing: -.03em; }
#hero p { color: var(--muted); margin-bottom: 0; }
#chat { border: 1px solid var(--line); border-radius: 18px; background: var(--paper); }
#prompt textarea, .message-wrap, .prose, .md { unicode-bidi: plaintext; text-align: start; }
#prompt textarea { direction: auto; font-size: 1rem; }
#run-button { min-height: 52px; }
.status { color: var(--muted); min-height: 28px; }
.privacy-note { color: var(--muted); font-size: .88rem; }
@media (max-width: 760px) {
#hero { padding: 19px; }
.gradio-container { padding: 9px !important; }
}
"""
THEME = gr.themes.Soft(
primary_hue="violet",
secondary_hue="cyan",
neutral_hue="slate",
)
with gr.Blocks(title="Muse Glimmer 30B", analytics_enabled=False) as demo:
selected_model = gr.Dropdown(
choices=MODEL_CHOICES,
value=MODEL_DEFAULT_ID,
label="Model checkpoint",
info="Choose the full BF16 or compact assistant checkpoint for this turn.",
interactive=True,
allow_custom_value=True,
)
model_history = gr.State([])
committed_chat = gr.State([])
selected_image = gr.State(None)
gr.Markdown(
"""
# Muse Glimmer · private inference
Text + image chat on either the official **full BF16** model or its **assistant checkpoint**.
Native greedy decoding is the default; Meta's sampling recipe is one click away. Reasoning is
parsed separately.
""",
elem_id="hero",
)
chatbot = gr.Chatbot(
label="Conversation",
height=570,
layout="panel",
buttons=["copy", "copy_all"],
reasoning_tags=[("<think>", "</think>")],
placeholder="Ask a question or attach an image to begin.",
sanitize_html=True,
elem_id="chat",
)
status = gr.Markdown(
"Ready · native greedy · reasoning high",
elem_classes="status",
)
with gr.Row(equal_height=True):
prompt = gr.Textbox(
label="Prompt",
placeholder="Ask in English, עברית, العربية, or another supported language…",
lines=3,
max_lines=9,
max_length=20_000,
autofocus=True,
scale=4,
elem_id="prompt",
)
image = gr.Image(
label="Optional image · this turn",
type="pil",
sources=["upload", "clipboard"],
height=180,
scale=2,
)
IMAGE_CHANGE_API_NAME = "set_image"
PRESET_CHANGE_API_NAME = "set_generation_preset"
STOP_API_NAME = "stop_generation"
CLEAR_API_NAME = "clear_conversation"
image.change(
_coerce_image_input,
inputs=image,
outputs=selected_image,
queue=False,
api_name=IMAGE_CHANGE_API_NAME,
api_visibility="private",
)
with gr.Row():
run_button = gr.Button("Generate", variant="primary", elem_id="run-button")
stop_button = gr.Button("Stop", variant="stop")
clear_button = gr.Button("Clear")
with gr.Accordion("Generation controls", open=False):
preset = gr.Radio(
choices=list(PRESETS),
value=NATIVE_GREEDY,
label="Preset",
info="Native greedy matches generation_config.json. Meta sampling applies the model-card recipe.",
)
with gr.Row():
reasoning_strength = gr.Dropdown(
choices=["low", "medium", "high", "xhigh"],
value="high",
label="Reasoning strength",
)
max_new_tokens = gr.Slider(
minimum=32,
maximum=MAX_NEW_TOKENS,
value=DEFAULT_MAX_NEW_TOKENS,
step=32,
label="Max new tokens",
info="App response budget; 512 is the default.",
)
repetition_penalty = gr.Slider(
minimum=0.8,
maximum=1.3,
value=DEFAULT_REPETITION_PENALTY,
step=0.01,
label="Repetition penalty",
)
do_sample = gr.Checkbox(
value=False,
label="Sampling",
info="Off is the checkpoint default. When off, temperature/top-p/top-k are ignored.",
)
with gr.Row():
temperature = gr.Slider(
minimum=0.05,
maximum=2.0,
value=DEFAULT_TEMPERATURE,
step=0.05,
label="Temperature",
)
top_p = gr.Slider(
minimum=0.05,
maximum=1.0,
value=DEFAULT_TOP_P,
step=0.01,
label="Top-p",
)
top_k = gr.Slider(
minimum=1,
maximum=200,
value=DEFAULT_TOP_K,
step=1,
label="Top-k",
)
with gr.Row():
seed = gr.Number(
value=DEFAULT_SEED,
precision=0,
minimum=0,
maximum=2_147_483_647,
label="Seed",
)
randomize_seed = gr.Checkbox(value=False, label="Randomize seed each turn")
show_reasoning = gr.Checkbox(value=True, label="Show reasoning")
system_prompt = gr.Textbox(
value="",
label="Optional system instruction",
placeholder="Blank uses the model's built-in helpful-assistant system message.",
lines=3,
max_length=20_000,
)
gr.Markdown(
f"""
**Private Space.** This app adds no prompt, reply, or image persistence and does not log
their contents. Inference runs on Hugging Face-hosted ZeroGPU `xlarge`; `xlarge` uses 2×
ZeroGPU quota. Model revisions: `{MODEL_REVISION}` and `{ASSISTANT_MODEL_REVISION}`.
No tools are connected or executed.
[Usage policy](https://huggingface.co/meta-models/Muse-Glimmer-30B/blob/{MODEL_REVISION}/USAGE_POLICY.md)
""",
elem_classes="privacy-note",
)
preset.change(
_set_preset,
inputs=preset,
outputs=[do_sample, temperature, top_p, top_k],
queue=False,
api_name=PRESET_CHANGE_API_NAME,
api_visibility="private",
)
generation_inputs = [
prompt,
selected_image,
selected_model,
chatbot,
model_history,
system_prompt,
reasoning_strength,
do_sample,
max_new_tokens,
temperature,
top_p,
top_k,
repetition_penalty,
seed,
randomize_seed,
show_reasoning,
]
generation_outputs = [chatbot, model_history, committed_chat, prompt, image, status]
generation_event = run_button.click(
fn=_generate_turn,
inputs=generation_inputs,
outputs=generation_outputs,
concurrency_limit=1,
concurrency_id="muse-glimmer-xlarge",
trigger_mode="once",
api_name="chat",
api_visibility="private",
api_description="Run a private Muse Glimmer text or image chat turn.",
show_progress="minimal",
validator=_validate_generation_request,
)
submit_event = prompt.submit(
fn=_generate_turn,
inputs=generation_inputs,
outputs=generation_outputs,
concurrency_limit=1,
concurrency_id="muse-glimmer-xlarge",
trigger_mode="once",
api_name=SUBMIT_API_NAME,
api_visibility="private",
api_description="Submit a private Muse Glimmer text or image chat turn.",
show_progress="minimal",
queue=True,
validator=_validate_generation_request,
)
stop_button.click(
_stop_conversation,
inputs=committed_chat,
outputs=[chatbot, status],
cancels=[generation_event, submit_event],
queue=False,
api_name=STOP_API_NAME,
api_visibility="private",
)
clear_button.click(
_clear_conversation,
inputs=None,
outputs=[chatbot, model_history, committed_chat, prompt, image, selected_image, status],
cancels=[generation_event, submit_event],
queue=False,
api_name=CLEAR_API_NAME,
api_visibility="private",
)
demo.queue(default_concurrency_limit=1, max_size=8)
if __name__ == "__main__":
demo.launch(theme=THEME, css=CSS)