Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
117 changes: 76 additions & 41 deletions SuperKittens/models/qwen/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,10 +3,25 @@
import json
from pathlib import Path

from .qwen import Qwen, Config
from .qwen import Qwen, Config, _load


_QWEN3_VARIANTS = {
"qwen3-8b": {
"hf_repo": "Qwen/Qwen3-8B",
"weight_dir": "Qwen3-8B-GGUF",
"gguf_name": "Qwen3-8B-Q8_0.gguf",
"default_quant": "q8_0",
"dims": dict(
n_layers=36, d_model=4096, n_heads=32, n_kv_heads=8, head_dim=128,
n_int=12288, vocab_size=151936, eps=1e-6, rope_freq_base=1_000_000.0,
tie_word_embeddings=0,
),
},
}


def _build_cfg_from_snapshot(snap: Path, **overrides) -> Config:
def _cfg_from_snapshot(snap: Path, **overrides) -> Config:
cfgj = json.loads((snap / "config.json").read_text())
cfg = Config(
n_layers = cfgj["num_hidden_layers"],
Expand All @@ -27,74 +42,94 @@ def _build_cfg_from_snapshot(snap: Path, **overrides) -> Config:
return cfg


_VARIANT_TO_DIR = {
"qwen3-0.6b": "Qwen3-0.6B",
}
_VARIANT_TO_GGUF = {
"qwen3-0.6b": "Qwen3-0.6B-Q8_0.gguf",
}
def _cfg_from_dims(variant: str, **overrides) -> Config:
dims = dict(_QWEN3_VARIANTS[variant]["dims"])
dims["seq_max"] = overrides.pop("seq_max", 128)
dims["cache_max"] = overrides.pop("cache_max", 512)
cfg = Config(**dims)
for k, v in overrides.items():
setattr(cfg, k, v)
return cfg


def _resolve_tokenizer(snap: Path, variant: str):
from SuperKittens.models.load.tokenizer import Tokenizer
for cand in (snap / "tokenizer.json", *snap.glob("tokenizer.json")):
if cand.exists():
return Tokenizer.from_hf_json(str(cand), family="qwen3")
repo = _QWEN3_VARIANTS[variant]["hf_repo"]
try:
from huggingface_hub import hf_hub_download
tok_path = hf_hub_download(repo_id=repo, filename="tokenizer.json")
return Tokenizer.from_hf_json(tok_path, family="qwen3")
except Exception as e:
print(f"[qwen] hf_hub_download tokenizer failed: {e}")
return None


def _from_pretrained(variant: str = "qwen3-0.6b", quant: str | None = None,
def _from_pretrained(variant: str = "qwen3-8b", quant: str | None = None,
snapshot: str | None = None, gguf_path: str | None = None,
**cfg_overrides) -> Qwen:
spec = variant
dir_name = _VARIANT_TO_DIR.get(spec.lower(), "Qwen3-0.6B")
spec = variant.lower()
if spec not in _QWEN3_VARIANTS:
raise ValueError(f"unknown qwen3 variant {spec!r}; known: {list(_QWEN3_VARIANTS)}")
meta = _QWEN3_VARIANTS[spec]

sk_root = Path(__file__).resolve().parents[3]
snap = Path(snapshot) if snapshot else (sk_root / "SuperKittens" / "model_weights" / dir_name)
if not snap.exists():
raise FileNotFoundError(f"snapshot dir not found: {snap}")
snap = Path(snapshot) if snapshot else (sk_root / "SuperKittens" / "model_weights" / meta["weight_dir"])

if (snap / "config.json").exists():
cfg = _cfg_from_snapshot(snap, **cfg_overrides)
else:
cfg = _cfg_from_dims(spec, **cfg_overrides)

cfg = _build_cfg_from_snapshot(snap, **cfg_overrides)
m = Qwen(cfg)
quant = quant or meta["default_quant"]

if quant in ("q8_0", "Q8_0", "gguf"):
gpath = Path(gguf_path) if gguf_path else (
sk_root / "SuperKittens" / "model_weights" / _VARIANT_TO_GGUF.get(spec.lower(), "Qwen3-0.6B-Q8_0.gguf"))
gpath = Path(gguf_path) if gguf_path else (snap / meta["gguf_name"])
if not gpath.exists():
raise FileNotFoundError(f"gguf file not found: {gpath}")
ggs = list(snap.glob("*.gguf")) if snap.exists() else []
if not ggs:
raise FileNotFoundError(f"gguf file not found: {gpath}")
gpath = ggs[0]
m.load_gguf(str(gpath))
else:
# fp16 safetensors path
from .qwen import _load
lib = _load() # ABI bound via QWEN_ABI; argtypes/restype set centrally
idx_path = snap / "model.safetensors.index.json"
single = snap / "model.safetensors"
target = idx_path if idx_path.exists() else single
if not target.exists():
if idx_path.exists():
if not hasattr(lib, "sk_qwen_load_safetensors_index"):
raise RuntimeError("libsk.dylib has no sk_qwen_load_safetensors_index symbol; rebuild dylib")
rc = lib.sk_qwen_load_safetensors_index(m._h, str(idx_path).encode())

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Bind safetensors-index ABI before calling via ctypes

When model.safetensors.index.json is present, this path calls lib.sk_qwen_load_safetensors_index(...) directly, but that symbol is not declared in QWEN_ABI (see SuperKittens/models/qwen/qwen.py), so bind() never assigns argtypes/restype. In ctypes, untyped calls default to int argument conversion, which can mis-handle the model handle pointer on 64-bit builds and cause load failures or crashes for sharded safetensor checkpoints. Add load_safetensors_index to QWEN_ABI (with the same signature as in C) before using this call path.

Useful? React with 👍 / 👎.

if rc:
raise RuntimeError(f"sk_qwen_load_safetensors_index failed: {rc}")
elif single.exists():
rc = lib.sk_qwen_load_safetensors(m._h, str(single).encode())
if rc:
raise RuntimeError(f"sk_qwen_load_safetensors failed: {rc}")
else:
raise FileNotFoundError(f"no safetensors in {snap}")
rc = _load().sk_qwen_load_safetensors(m._h, str(target).encode())
if rc:
raise RuntimeError(f"sk_qwen_load_safetensors failed: {rc}")

# RoPE tables
m.bake_and_set_rope()

# Tokenizer
try:
from SuperKittens.models.load.tokenizer import Tokenizer
json_path = snap / "tokenizer.json"
sp_path = snap / "tokenizer.model"
if json_path.exists():
m.tokenizer = Tokenizer.from_hf_json(str(json_path), family="qwen3")
elif sp_path.exists():
m.tokenizer = Tokenizer.from_sentencepiece(str(sp_path))
else:
print(f"[qwen] no tokenizer.json or tokenizer.model in {snap}")
except Exception as e:
print(f"[qwen] tokenizer attach failed: {e}")
tok = _resolve_tokenizer(snap, spec)
if tok is not None:
m.tokenizer = tok
else:
print(f"[qwen] no tokenizer available for {spec}")

return m


# Adapter so MODEL_REGISTRY's cls.from_pretrained(**defaults, **kwargs) routes here.
class _QwenFactory:
@staticmethod
def from_pretrained(**kwargs):
return _from_pretrained(**kwargs)


from SuperKittens.api import register
for _spec in ("qwen3-0.6b",):
for _spec in _QWEN3_VARIANTS:
try:
register(_spec, _QwenFactory, variant=_spec)
except ValueError:
Expand Down
5 changes: 3 additions & 2 deletions SuperKittens/models/qwen/qwen.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,8 +67,9 @@ class _Weights(ctypes.Structure):
ctypes.c_uint32, ctypes.POINTER(ctypes.c_int32)], ctypes.c_int),
"reset": ([ctypes.c_void_p], None),
"destroy": ([ctypes.c_void_p], None),
"load_safetensors": ([ctypes.c_void_p, ctypes.c_char_p], ctypes.c_int),
"load_gguf": optional([ctypes.c_void_p, ctypes.c_char_p], ctypes.c_int),
"load_safetensors": ([ctypes.c_void_p, ctypes.c_char_p], ctypes.c_int),
"load_safetensors_index": optional([ctypes.c_void_p, ctypes.c_char_p], ctypes.c_int),
"load_gguf": optional([ctypes.c_void_p, ctypes.c_char_p], ctypes.c_int),
"set_rope_tables": optional([ctypes.c_void_p, ctypes.c_void_p, ctypes.c_void_p], ctypes.c_int),
"get_last_logits": optional([ctypes.c_void_p, ctypes.c_void_p], ctypes.c_int),
}
Expand Down