"""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()