"""Chat with Whittle-Qwen-3.8-45B-A3B on ZeroGPU. The model repo is mounted read-only at /models/w45 (no 90 GB download into the Space); small files come from the Hub. The 35 B body runs on a full RTX PRO 6000 (xlarge); the 10 B n-gram table stays in CPU memory, as with llama.cpp's `-ot per_layer_token_embd=CPU`. """ import os import shutil import threading import time import spaces # must come before torch/CUDA use on ZeroGPU import gradio as gr import torch from transformers import AutoTokenizer, TextIteratorStreamer from whittle_load import MmapNGramTable, load_whittle def _memory_logger(every: float = 10.0): """Log the container's memory (cgroup) every few seconds while the model loads: anon = tensors, file = page cache, shmem = shared memory. Diagnoses the 104 GB limit.""" import time def read(): stat = {} try: for line in open("/sys/fs/cgroup/memory.stat"): k, v = line.split() stat[k] = int(v) cur = int(open("/sys/fs/cgroup/memory.current").read()) except Exception: return "cgroup memory stats unavailable" g = lambda k: stat.get(k, 0) / 2**30 return f"current {cur / 2**30:.1f}G | anon {g('anon'):.1f}G | file {g('file'):.1f}G | shmem {g('shmem'):.1f}G | file_mapped {g('file_mapped'):.1f}G" def loop(): while True: print(f"[mem] {read()}", flush=True) time.sleep(every) threading.Thread(target=loop, daemon=True).start() _memory_logger() REPO = "logic65/Whittle-Qwen-3.8-45B-A3B" MOUNT = "/models/w45" LOCAL = "/tmp/w45" def prepare_local_dir() -> str: """Small files (config, index, tokenizer, template) come straight from the Hub; the weight shards are read from the read-only mount through symlinks, after checking every shard's size on the mount against the Hub.""" from huggingface_hub import HfApi, hf_hub_download files = {f.path: f.size for f in HfApi().list_repo_tree(REPO) if "/" not in f.path and getattr(f, "size", None) is not None} os.makedirs(LOCAL, exist_ok=True) for name, size in sorted(files.items()): dst = os.path.join(LOCAL, name) if name.endswith(".safetensors"): src = os.path.join(MOUNT, name) got = os.path.getsize(src) if os.path.exists(src) else None print(f"[mount] {name}: Hub {size:,} B, mount {got if got is None else f'{got:,} B'}", flush=True) if got != size: raise RuntimeError(f"{name} on the mount does not match the Hub ({got} vs {size} bytes)") if not os.path.lexists(dst): os.symlink(src, dst) elif not name.endswith((".md", ".svg", ".py")) and name != ".gitattributes": shutil.copy(hf_hub_download(REPO, name), dst) # transformers gets an index WITHOUT the 20 GB table pieces, so it never reads them while loading (the table is # memory-mapped by whittle_load from the full index, kept beside it). import json full = os.path.join(LOCAL, "model.safetensors.index.json") idx = json.load(open(full)) shutil.copy(full, os.path.join(LOCAL, "whittle_full_index.json")) idx["weight_map"] = {k: v for k, v in idx["weight_map"].items() if not k.startswith("model.ngram_embedding.shard_")} json.dump(idx, open(full, "w")) return LOCAL SOURCE = prepare_local_dir() tokenizer = AutoTokenizer.from_pretrained(SOURCE) # table="mmap": the 20 GB n-gram table is read in place from the mounted shards (only looked-up rows), which keeps the # Space under ZeroGPU's 104 GB host-memory limit (the 70 GB body is held in host memory until a GPU attaches). model = load_whittle(SOURCE, dtype=torch.bfloat16, device_map="cuda", table="mmap").eval() PLE = model.get_submodule(f"model.layers.{model.config.ple_layer_ids[0] - 1}.ple.ple_embedding") def _table_to_ram(): """Once ZeroGPU has packed the weights out of this container (memory falls from ~90 GB to ~6 GB), read the 20 GB table into RAM: lookups then hit memory instead of the network mount. Forked GPU workers share it copy-on-write.""" deadline = time.time() + 1800 while time.time() < deadline: try: if int(open("/sys/fs/cgroup/memory.current").read()) < 40 * 2**30: break except Exception: break time.sleep(10) if not isinstance(PLE.ngram_embedding, MmapNGramTable): return t0 = time.time() print("[table] reading the n-gram table into RAM (64 MB sequential reads)", flush=True) try: emb = PLE.ngram_embedding.to_ram() except Exception as e: print(f"[table] staying memory-mapped: {e!r}", flush=True) return PLE.ngram_embedding = emb dt = time.time() - t0 print(f"[table] n-gram table in RAM: {tuple(emb.weight.shape)} in {dt:.0f}s ({emb.weight.numel() * emb.weight.element_size() / 2**30 / dt:.2f} GiB/s)", flush=True) threading.Thread(target=_table_to_ram, daemon=True).start() def _gpu_seconds(message, history, thinking, max_new_tokens, temperature): # ZeroGPU checks the REQUESTED duration (x2 on the full-size GPU) against the visitor's remaining daily quota before # starting: a free account has 300 s, so the request must stay <= 150 s or free visitors can never run it. return int(max(60, min(140, 45 + max_new_tokens / 10))) def _render(text: str, thinking: bool): text = text.removeprefix("").lstrip("\n") if thinking else text if thinking: reasoning, sep, answer = text.partition("") if not sep: return [gr.ChatMessage(role="assistant", content=reasoning, metadata={"title": "Thinking", "status": "pending"})] return [gr.ChatMessage(role="assistant", content=reasoning.strip(), metadata={"title": "Thinking", "status": "done"}), gr.ChatMessage(role="assistant", content=answer.strip())] return text.replace("", "").replace("", "").strip() @spaces.GPU(size="xlarge", duration=_gpu_seconds) def respond(message, history, thinking, max_new_tokens, temperature): t_gpu = time.time() if model.config._experts_implementation != "grouped_mm": # the load-time check ran without a GPU and chose "eager" try: model.set_experts_implementation("grouped_mm") except Exception as e: print(f"[timing] grouped_mm unavailable ({e!r}); trying batched_mm", flush=True) model.set_experts_implementation("batched_mm") messages = [{"role": m["role"], "content": m["content"]} for m in history if m.get("role") in ("user", "assistant") and isinstance(m.get("content"), str) and not (m.get("metadata") or {}).get("title")] messages.append({"role": "user", "content": message}) inputs = tokenizer.apply_chat_template(messages, add_generation_prompt=True, return_dict=True, return_tensors="pt", enable_thinking=bool(thinking)).to("cuda") streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=300) generation = dict(**inputs, streamer=streamer, max_new_tokens=int(max_new_tokens), do_sample=True, temperature=float(temperature), top_p=0.8, top_k=20, repetition_penalty=1.05) threading.Thread(target=model.generate, kwargs=generation, daemon=True).start() text, first, n_pieces = "", None, 0 for piece in streamer: if first is None: first = time.time() text += piece n_pieces += 1 yield _render(text, bool(thinking)) done = time.time() n_tok = len(tokenizer(text, add_special_tokens=False)["input_ids"]) print(f"[timing] experts={model.config._experts_implementation} table={type(PLE.ngram_embedding).__name__} prompt={inputs['input_ids'].shape[1]} tok | " f"first token {((first or done) - t_gpu):.1f}s after GPU attach | {n_tok} tok in {done - (first or done):.1f}s " f"({n_tok / max(done - (first or done), 1e-6):.1f} tok/s) | total {done - t_gpu:.1f}s", flush=True) DESCRIPTION = """**Whittle-Qwen-3.8-45B-A3B**: a 45 B-parameter, ~3 B-active mixture of experts with a 10 B n-gram memory, running on stock `transformers`. Experimental: the flagship 35B plus the 76 experts per layer its prune removed, put back without retraining. Weights: [safetensors](https://huggingface.co/logic65/Whittle-Qwen-3.8-45B-A3B) ยท [GGUF for llama.cpp](https://huggingface.co/logic65/Whittle-Qwen-3.8-45B-A3B-GGUF) (much faster locally). GPU time on this demo comes out of your own daily ZeroGPU quota (5 min for a free account, at double cost on this full-size GPU). Thinking is off by default; switch it on for maths and code. Built by one person on a grocery budget: [ko-fi.com/davida81328](https://ko-fi.com/davida81328).""" demo = gr.ChatInterface( respond, title="Whittle 45B", description=DESCRIPTION, additional_inputs=[ gr.Checkbox(False, label="Thinking (slower; better on maths and code)"), gr.Slider(64, 1024, value=384, step=64, label="Max new tokens"), gr.Slider(0.1, 1.2, value=0.7, step=0.05, label="Temperature"), ], examples=[["Who wrote the novel The Remains of the Day, and what is it about?"], ["Write a Python function that checks whether a string is a palindrome, ignoring punctuation."], ["A train leaves at 14:35 and arrives at 17:10. How long is the journey?"]], cache_examples=False, ) if __name__ == "__main__": demo.launch()