Spaces:
Running on Zero
Running on Zero
Download app.py from khtsly/Luau-Coder-Preview: direct link, hf CLI and curl.
- Browser
- Download file 11.1 kB
-
https://huggingface.co/spaces/khtsly/Luau-Coder-Preview/resolve/main/app.py
- Command line
-
hf download hf://spaces/khtsly/Luau-Coder-Preview/app.py
-
curl -L -o app.py https://huggingface.co/spaces/khtsly/Luau-Coder-Preview/resolve/main/app.py
11.1 kB
| import gc | |
| import os | |
| import tempfile | |
| from collections.abc import Generator | |
| from functools import wraps | |
| from pathlib import Path | |
| from threading import Lock, Thread | |
| import spaces | |
| import torch | |
| import transformers | |
| from fastapi.responses import HTMLResponse | |
| from gradio import Server | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, LogitsProcessor, LogitsProcessorList, TextIteratorStreamer | |
| print(f"Transformers version: {transformers.__version__}") | |
| HF_TOKEN = os.environ.get("HF_TOKEN") | |
| # Available model variants. The front-end populates its selector from the | |
| # /models endpoint, so adding a new entry here is enough to offer it. | |
| # Only one variant is kept loaded at a time; switching unloads the previous. | |
| MODEL_VARIANTS = { | |
| "Nalu 1.0 (Base)": {"repo_id": "khtsly/Nalu-1.0-Base", "tokenizer": "khtsly/luau-coder-1.0-preview-tokenizer", "enabled": True}, | |
| "Moana 1.0 (Base)": {"repo_id": "khtsly/Moana-1.0-Base", "tokenizer": "khtsly/luau-coder-1.0-preview-tokenizer", "enabled": True}, | |
| "Moana 1.5": {"repo_id": "khtsly/Moana-1.5-Base", "tokenizer": "khtsly/luau-coder-1.5-tokenizer", "enabled": False}, | |
| } | |
| DEFAULT_VARIANT = "Nalu 1.0 (Base)" | |
| def _repair_checkpoint_keys(weights_path: Path) -> bool: | |
| """Remove an obsolete composite-model prefix before Transformers loads.""" | |
| from safetensors import safe_open | |
| with safe_open(weights_path, framework="pt", device="cpu") as f: | |
| keys = list(f.keys()) | |
| prefixed = [key for key in keys if key.startswith("language_model.")] | |
| if not prefixed: | |
| return False | |
| if len(prefixed) != len(keys): | |
| raise RuntimeError( | |
| f"Refusing to repair mixed checkpoint keys in {weights_path}: " | |
| f"{len(prefixed)} of {len(keys)} use the language_model. prefix" | |
| ) | |
| fixed_keys = [key.removeprefix("language_model.") for key in keys] | |
| if len(set(fixed_keys)) != len(fixed_keys): | |
| raise RuntimeError(f"Prefix removal would create duplicate keys in {weights_path}") | |
| from safetensors.torch import load_file, save_file | |
| print(f"Repairing {len(keys)} checkpoint keys in {weights_path} ...") | |
| state = load_file(weights_path, device="cpu") | |
| fixed = { | |
| key.removeprefix("language_model."): tensor.contiguous() | |
| for key, tensor in state.items() | |
| } | |
| tmp_path = weights_path.with_name(f".{weights_path.name}.repair.tmp") | |
| try: | |
| save_file(fixed, tmp_path) | |
| os.replace(tmp_path, weights_path) | |
| finally: | |
| tmp_path.unlink(missing_ok=True) | |
| print("Checkpoint key repair complete") | |
| return True | |
| def _prepare_model_source(variant: str) -> str: | |
| """Return a local model directory whose safetensors keys are loadable.""" | |
| repo_id = MODEL_VARIANTS[variant]["repo_id"] | |
| cache_dir = Path(tempfile.gettempdir()) / f"model-{variant}" | |
| # Local override for the default variant (checked-in model.safetensors). | |
| if variant == DEFAULT_VARIANT: | |
| app_dir = Path(__file__).resolve().parent | |
| weights_path = app_dir / "model.safetensors" | |
| if weights_path.is_file(): | |
| _repair_checkpoint_keys(weights_path) | |
| return str(app_dir) | |
| from huggingface_hub import snapshot_download | |
| snapshot_download( | |
| repo_id=repo_id, | |
| local_dir=cache_dir, | |
| token=HF_TOKEN, | |
| allow_patterns=["*.json", "*.py", "*.jinja", "*.safetensors"], | |
| ) | |
| weights_path = cache_dir / "model.safetensors" | |
| if not weights_path.is_file(): | |
| raise FileNotFoundError(f"No model.safetensors found in {cache_dir}") | |
| _repair_checkpoint_keys(weights_path) | |
| return str(cache_dir) | |
| _model = None | |
| _tokenizer = None | |
| _current_variant = None | |
| _model_lock = Lock() | |
| def _unload_current_model() -> None: | |
| """Free the currently loaded variant so the next one can take its place.""" | |
| global _model, _current_variant, _tokenizer | |
| if _model is not None: | |
| print(f"Unloading model variant: {_current_variant}") | |
| del _model | |
| _model = None | |
| if _tokenizer is not None: | |
| print(f"Unloading tokenizer variant: {_current_variant}") | |
| del _tokenizer | |
| _tokenizer = None | |
| _current_variant = None | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| def ensure_model(variant: str): | |
| """Load `variant` on demand, unloading whatever is currently loaded.""" | |
| global _model, _current_variant, _tokenizer | |
| if variant not in MODEL_VARIANTS: | |
| print(f"Unknown model variant {variant!r}; falling back to {DEFAULT_VARIANT}") | |
| variant = DEFAULT_VARIANT | |
| if not MODEL_VARIANTS[variant].get("enabled", True): | |
| print(f"Model variant {variant!r} is not enabled yet; falling back to {DEFAULT_VARIANT}") | |
| variant = DEFAULT_VARIANT | |
| with _model_lock: | |
| if _model is not None and _current_variant == variant: | |
| return _model, _tokenizer | |
| _unload_current_model() | |
| model_targeted = MODEL_VARIANTS[variant] | |
| print(f"Loading model variant: {variant} ({model_targeted["repo_id"]}) | tokenizer: {model_targeted["tokenizer"]} ...") | |
| source = _prepare_model_source(variant) | |
| loaded = AutoModelForCausalLM.from_pretrained( | |
| source, | |
| trust_remote_code=True, | |
| dtype=torch.bfloat16, | |
| token=HF_TOKEN, | |
| device_map="cuda", | |
| ) | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| model_targeted["tokenizer"], | |
| trust_remote_code=True, | |
| token=HF_TOKEN, | |
| ) | |
| _install_transformers_516_cache_compat(loaded) | |
| _model = loaded | |
| _current_variant = variant | |
| _tokenizer = tokenizer | |
| print(f"Model variant ready: {variant}") | |
| return _model, tokenizer | |
| def _install_transformers_516_cache_compat(loaded_model) -> None: | |
| """Handle empty generation positions with older Kimi remote code.""" | |
| inner_model = loaded_model.model | |
| original_forward = inner_model.forward | |
| def compatible_forward(*args, **kwargs): | |
| for name in ("cache_position", "position_ids"): | |
| value = kwargs.get(name) | |
| numel = getattr(value, "numel", None) | |
| if callable(numel) and numel() == 0: | |
| kwargs[name] = None | |
| return original_forward(*args, **kwargs) | |
| inner_model.forward = compatible_forward | |
| class PresencePenaltyLogitsProcessor(LogitsProcessor): | |
| """OpenAI-style presence penalty for HF generate (which lacks it natively). | |
| Subtracts `penalty` once from the logit of every unique token already | |
| present in the context, discouraging the model from reusing tokens. | |
| """ | |
| def __init__(self, penalty: float): | |
| self.penalty = float(penalty) | |
| def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor: | |
| if self.penalty == 0: | |
| return scores | |
| for batch_idx, ids in enumerate(input_ids): | |
| for token_id in set(ids.tolist()): | |
| if 0 <= token_id < scores.shape[-1]: | |
| scores[batch_idx, token_id] -= self.penalty | |
| return scores | |
| demo = Server() | |
| def predict( | |
| message: str, | |
| history: list[list] | None = None, | |
| max_tokens: int = 32, | |
| temperature: float = 0.6, | |
| top_p: float = 0.95, | |
| top_k: int = 40, | |
| min_p: float = 0.05, | |
| repetition_penalty: float = 1.05, | |
| presence_penalty: float = 0.0, | |
| model_id: str = "Nalu 1.0 (Base)", | |
| ) -> Generator[str, None, None]: | |
| # History is accepted for front-end chat compatibility; generation uses | |
| # the current message only to preserve base-model continuation behavior. | |
| raw_input_text = message | |
| max_tokens = int(max_tokens) | |
| active_model, tokenizer = ensure_model(model_id) | |
| if next(active_model.parameters()).device.type != "cuda": | |
| active_model.to("cuda") | |
| inputs = tokenizer(raw_input_text, return_tensors="pt") | |
| token_count = inputs["input_ids"].shape[1] | |
| print(f"Generation request: {len(raw_input_text)} chars, {token_count} tokens") | |
| if token_count == 0: | |
| seed_token_id = tokenizer.bos_token_id | |
| if seed_token_id is None: | |
| candidate = tokenizer.convert_tokens_to_ids("[BOS]") | |
| if isinstance(candidate, int) and candidate >= 0: | |
| seed_token_id = candidate | |
| if seed_token_id is None: | |
| yield "Please enter a non-empty prompt." | |
| return | |
| inputs["input_ids"] = torch.tensor([[seed_token_id]], dtype=torch.long) | |
| inputs["attention_mask"] = torch.ones((1, 1), dtype=torch.long) | |
| token_count = 1 | |
| print(f"Empty prompt: seeded generation with token {seed_token_id}") | |
| inputs = inputs.to("cuda") | |
| streamer = TextIteratorStreamer(tokenizer, skip_prompt=True) | |
| generation_kwargs = dict( | |
| **inputs, | |
| streamer=streamer, | |
| max_new_tokens=int(max_tokens), | |
| pad_token_id=tokenizer.eos_token_id, | |
| use_cache=True, | |
| generation_mode=True, | |
| ) | |
| temperature = float(temperature) | |
| top_p = float(top_p) | |
| top_k = int(top_k) | |
| min_p = float(min_p) | |
| repetition_penalty = float(repetition_penalty) | |
| presence_penalty = float(presence_penalty) | |
| if temperature > 0.0: | |
| generation_kwargs.update(do_sample=True, temperature=temperature) | |
| if 0.0 < top_p < 1.0: | |
| generation_kwargs["top_p"] = top_p | |
| if top_k > 0: | |
| generation_kwargs["top_k"] = top_k | |
| if 0.0 < min_p <= 1.0: | |
| generation_kwargs["min_p"] = min_p | |
| else: | |
| generation_kwargs.update(do_sample=False) | |
| if repetition_penalty != 1.0: | |
| generation_kwargs["repetition_penalty"] = repetition_penalty | |
| if presence_penalty != 0.0: | |
| generation_kwargs["logits_processor"] = LogitsProcessorList( | |
| [PresencePenaltyLogitsProcessor(presence_penalty)] | |
| ) | |
| generation_error = [] | |
| def run_generation(): | |
| try: | |
| active_model.generate(**generation_kwargs) | |
| except Exception as exc: | |
| generation_error.append(exc) | |
| streamer.end() | |
| thread = Thread(target=run_generation) | |
| thread.start() | |
| output_buffer = "" | |
| for new_text in streamer: | |
| output_buffer += new_text | |
| yield output_buffer | |
| thread.join() | |
| if generation_error: | |
| exc = generation_error[0] | |
| message_out = f"Generation failed: {type(exc).__name__}: {exc}" | |
| print(message_out) | |
| yield f"{output_buffer}\n\n{message_out}" if output_buffer else message_out | |
| async def list_models(): | |
| return { | |
| "default": DEFAULT_VARIANT, | |
| "variants": [ | |
| {"id": name, "enabled": bool(cfg.get("enabled", True))} | |
| for name, cfg in sorted(MODEL_VARIANTS.items()) | |
| ], | |
| } | |
| async def homepage(): | |
| html_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "index.html") | |
| with open(html_path, "r", encoding="utf-8") as f: | |
| return f.read() | |
| if __name__ == "__main__": | |
| demo.launch(show_error=True) |