Merge pull request #257 from woolcoxm/windows-optimizations
Windows disk I/O: pread + PIPE + compat_fadvise (1.70s/tok, mmap reverted)
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
@echo off
|
||||
call "C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools\VC\Auxiliary\Build\vcvars64.bat"
|
||||
cd /d C:\Users\Mark\Desktop\Projects\colibri\c
|
||||
"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.8\bin\nvcc" -O3 -std=c++17 -arch=sm_120 -Xcompiler=-W3 -shared -DCOLI_CUDA_BUILDING_DLL -L"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.8\lib/x64" -lcudart backend_cuda.cu -o coli_cuda.dll
|
||||
echo EXITCODE=%ERRORLEVEL%
|
||||
+35
-2
@@ -95,7 +95,18 @@ static inline int compat_open_direct(const char *path){
|
||||
* prevents 0x0A bytes from being silently translated to \r\n. */
|
||||
#define COMPAT_O_RDONLY (O_RDONLY | O_BINARY)
|
||||
|
||||
/* --- posix_fadvise: no-op (advisory only; safe to ignore) --- */
|
||||
/* --- posix_fadvise: Windows has no direct equivalent. Semantics:
|
||||
* WILLNEED -> warm the OS page cache so a later synchronous pread finds the
|
||||
* pages resident. Implemented as an overlapped background ReadFile
|
||||
* into a throwaway scratch buffer (fire-and-forget readahead). Called
|
||||
* from the dedicated PILOT I/O thread / next-block readahead in moe(),
|
||||
* NEVER inline on the hot path (the existing comment at glm.c:2847
|
||||
* measures inline fadvise submit at ~0.5ms x 169k calls = +92s/48tok).
|
||||
* Each call owns its OVERLAPPED + scratch buffer -> thread-safe.
|
||||
* DONTNEED -> no-op: Windows' standby-list trimming self-regulates under pressure,
|
||||
* and on a low-RAM host keeping the pages is what we want for reuse.
|
||||
* Matches macOS (compat.h:16-19) which no-ops DONTNEED for the same
|
||||
* reason. The engine only ever uses DONTNEED as an advisory. */
|
||||
#ifndef POSIX_FADV_NORMAL
|
||||
#define POSIX_FADV_NORMAL 0
|
||||
#define POSIX_FADV_RANDOM 1
|
||||
@@ -104,7 +115,29 @@ static inline int compat_open_direct(const char *path){
|
||||
#define POSIX_FADV_DONTNEED 4
|
||||
#define POSIX_FADV_NOREUSE 5
|
||||
#endif
|
||||
#define posix_fadvise(fd,off,len,advice) do{(void)(fd);(void)(off);(void)(len);(void)(advice);}while(0)
|
||||
static inline int compat_fadvise(int fd, off_t off, off_t len, int advice){
|
||||
if(advice!=POSIX_FADV_WILLNEED || len<=0) return 0;
|
||||
intptr_t osfh=_get_osfhandle(fd);
|
||||
if(osfh==-1 || osfh==-2) return 0;
|
||||
HANDLE h=(HANDLE)osfh;
|
||||
/* Cap the readahead window: reading a whole 19MB expert per hint is fine on the
|
||||
* PILOT thread, but a pathological huge len would spike transient memory. */
|
||||
size_t rdlen = (len>(off_t)(64*1024*1024)) ? (size_t)(64*1024*1024) : (size_t)len;
|
||||
char *buf=(char*)_aligned_malloc(rdlen, 4096);
|
||||
if(!buf) return -1;
|
||||
OVERLAPPED ov={0};
|
||||
ov.Offset = (DWORD)( (off_t)off & 0xFFFFFFFFULL);
|
||||
ov.OffsetHigh = (DWORD)(((off_t)off >> 32) & 0xFFFFFFFFULL);
|
||||
/* Issue an overlapped read. With a non-OVERLAPPED-opened handle ReadFile still
|
||||
* accepts lpOverlapped (it carries the 64-bit offset) and blocks until the read
|
||||
* completes — but crucially it populates the standby page cache for this region,
|
||||
* so the later synchronous pread on the same offsets faults from RAM not disk. */
|
||||
DWORD got=0;
|
||||
ReadFile(h, buf, (DWORD)rdlen, &got, &ov);
|
||||
_aligned_free(buf);
|
||||
return 0;
|
||||
}
|
||||
#define posix_fadvise compat_fadvise
|
||||
|
||||
/* --- pread -> ReadFile + OVERLAPPED su raw OS handle ---
|
||||
* Thread-safe (no shared seek position). Gestisce offset >4 GB e chunking
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Download GLM-5.2-FP8 from ModelScope (fast, no HF throttling) with
|
||||
HuggingFace fallback. Parallel shard download, clean progress display.
|
||||
Usage: python download_fp8.py
|
||||
python download_fp8.py --parallel 4
|
||||
python download_fp8.py --source hf (force HuggingFace)
|
||||
"""
|
||||
import os, sys, time, threading, argparse, subprocess
|
||||
|
||||
REPO_MS = "ZhipuAI/GLM-5.2-FP8" # ModelScope
|
||||
REPO_HF = "zai-org/GLM-5.2-FP8" # HuggingFace
|
||||
DEST = r"I:\glm52_fp8"
|
||||
|
||||
# ── ANSI colors ──
|
||||
class C:
|
||||
dim="\033[2m"; grn="\033[32m"; yel="\033[33m"; cyn="\033[36m"; b="\033[1m"; r="\033[0m"
|
||||
|
||||
def fmt_bytes(n):
|
||||
if n>=1e9: return f"{n/1e9:.2f} GB"
|
||||
if n>=1e6: return f"{n/1e6:.1f} MB"
|
||||
return f"{n/1e3:.0f} KB"
|
||||
|
||||
def fmt_time(s):
|
||||
if s<60: return f"{s:.0f}s"
|
||||
if s<3600: return f"{s/60:.0f}m"
|
||||
return f"{s/3600:.1f}h"
|
||||
|
||||
def bar(cur, total, width=24):
|
||||
if total<=0: return "["+" "*width+"]"
|
||||
pct=min(cur/total,1.0); filled=int(width*pct)
|
||||
return "["+"█"*filled+"░"*(width-filled)+f"] {pct*100:4.0f}%"
|
||||
|
||||
def get_shard_list_hf():
|
||||
from huggingface_hub import HfApi
|
||||
info=HfApi().repo_info(REPO_HF, files_metadata=True)
|
||||
shards=sorted(s.rfilename for s in info.siblings if s.rfilename.endswith(".safetensors"))
|
||||
sizes={s.rfilename:s.size for s in info.siblings if s.rfilename.endswith(".safetensors")}
|
||||
return shards, sizes
|
||||
|
||||
def get_shard_list_ms():
|
||||
"""Get shard list from ModelScope API."""
|
||||
import requests
|
||||
# ModelScope API: list files
|
||||
r = requests.get(f"https://modelscope.cn/api/v1/models/{REPO_MS}/repo/files?Revision=master&Root=", timeout=30)
|
||||
data = r.json()["Data"]["Files"]
|
||||
shards = sorted(f["Path"] for f in data if f["Path"].endswith(".safetensors"))
|
||||
sizes = {f["Path"]: f.get("Size", 0) for f in data if f["Path"].endswith(".safetensors")}
|
||||
return shards, sizes
|
||||
|
||||
def download_file_ms(fn):
|
||||
"""Download a single file from ModelScope using their CDN."""
|
||||
from modelscope.hub.file_download import model_file_download
|
||||
model_file_download(
|
||||
model_id=REPO_MS,
|
||||
file_path=fn,
|
||||
local_dir=DEST,
|
||||
revision="master",
|
||||
)
|
||||
|
||||
def download_file_hf(fn):
|
||||
"""Download a single file from HuggingFace with hf_transfer."""
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
from huggingface_hub import hf_hub_download
|
||||
hf_hub_download(REPO_HF, fn, local_dir=DEST)
|
||||
|
||||
def download_file_curl(fn, base_url, expected_size):
|
||||
"""Fallback: download with curl to a .part file with resume."""
|
||||
outpath = os.path.join(DEST, fn)
|
||||
partpath = outpath + ".part"
|
||||
if os.path.exists(outpath) and os.path.getsize(outpath) == expected_size:
|
||||
return True
|
||||
url = f"{base_url}/{fn}"
|
||||
cmd = ["curl", "-L", "-C", "-", "--retry", "999", "--retry-delay", "5",
|
||||
"--connect-timeout", "15", "--speed-time", "30", "--speed-limit", "1000",
|
||||
"-o", partpath, "-H", "User-Agent: colibri-download/1.0", url]
|
||||
subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
||||
if os.path.exists(partpath) and os.path.getsize(partpath) >= expected_size:
|
||||
os.replace(partpath, outpath)
|
||||
return True
|
||||
return False
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--parallel", type=int, default=3)
|
||||
ap.add_argument("--source", choices=["auto", "ms", "hf"], default="auto",
|
||||
help="auto (try ModelScope first), ms (ModelScope only), hf (HuggingFace only)")
|
||||
args = ap.parse_args()
|
||||
|
||||
os.makedirs(DEST, exist_ok=True)
|
||||
|
||||
# Determine source and get shard list
|
||||
use_ms = False
|
||||
shards, sizes = [], {}
|
||||
|
||||
if args.source in ("auto", "ms"):
|
||||
try:
|
||||
print(f"{C.dim}Trying ModelScope...{C.r}", end=" ", flush=True)
|
||||
shards, sizes = get_shard_list_ms()
|
||||
use_ms = True
|
||||
print(f"{C.grn}✓{C.r} {len(shards)} shards found")
|
||||
except Exception as e:
|
||||
print(f"{C.yel}failed ({e}){C.r}")
|
||||
if args.source == "ms":
|
||||
print("ModelScope failed and --source ms was set. Exiting."); return
|
||||
|
||||
if not shards:
|
||||
print(f"{C.dim}Using HuggingFace...{C.r}", end=" ", flush=True)
|
||||
shards, sizes = get_shard_list_hf()
|
||||
print(f"{C.grn}✓{C.r} {len(shards)} shards found")
|
||||
|
||||
total = len(shards)
|
||||
total_bytes = sum(sizes.values())
|
||||
source_name = "ModelScope" if use_ms else "HuggingFace"
|
||||
|
||||
# Download metadata files
|
||||
meta_files = ["config.json", "tokenizer.json", "tokenizer_config.json",
|
||||
"generation_config.json", "model.safetensors.index.json"]
|
||||
for fn in meta_files:
|
||||
out = os.path.join(DEST, fn)
|
||||
if not os.path.exists(out):
|
||||
try:
|
||||
if use_ms: download_file_ms(fn)
|
||||
else: download_file_hf(fn)
|
||||
except Exception: pass
|
||||
|
||||
# Build work queue
|
||||
todo = []
|
||||
done_set = set()
|
||||
for fn in shards:
|
||||
outpath = os.path.join(DEST, fn)
|
||||
if os.path.exists(outpath) and os.path.getsize(outpath) == sizes.get(fn, 0):
|
||||
done_set.add(fn)
|
||||
else:
|
||||
todo.append(fn)
|
||||
|
||||
existing = len(done_set)
|
||||
print(f"\n{C.b}GLM-5.2-FP8 Download ({source_name}){C.r}")
|
||||
print(f" {C.dim}{total} shards · {total_bytes/1e9:.0f} GB · {existing}/{total} complete{C.r}")
|
||||
print(f" {C.dim}{len(todo)} to download · {args.parallel} parallel{C.r}")
|
||||
if todo:
|
||||
remaining = sum(sizes[fn] for fn in todo)
|
||||
print(f" {C.dim}Remaining: {remaining/1e9:.0f} GB{C.r}")
|
||||
print()
|
||||
|
||||
if not todo:
|
||||
print(f"{C.grn}✓ All shards already downloaded!{C.r}\n"); return
|
||||
|
||||
lock = threading.Lock()
|
||||
completed = list(done_set)
|
||||
t0 = time.time()
|
||||
qidx = [0]
|
||||
|
||||
def worker(wid):
|
||||
while True:
|
||||
with lock:
|
||||
if qidx[0] >= len(todo): return
|
||||
idx = qidx[0]; qidx[0] += 1
|
||||
fn = todo[idx]
|
||||
|
||||
expected = sizes.get(fn, 0)
|
||||
shard_num = existing + idx + 1
|
||||
print(f" {C.cyn}[{shard_num}/{total}]{C.r} {C.dim}{fn}{C.r}")
|
||||
|
||||
success = False
|
||||
for attempt in range(3):
|
||||
try:
|
||||
if use_ms:
|
||||
download_file_ms(fn)
|
||||
else:
|
||||
download_file_hf(fn)
|
||||
outpath = os.path.join(DEST, fn)
|
||||
# ModelScope downloads to a cache dir, need to check
|
||||
# if the file exists at our expected path
|
||||
if not os.path.exists(outpath):
|
||||
# Try to find it in ModelScope's cache structure
|
||||
# and copy/symlink it
|
||||
pass
|
||||
if os.path.exists(outpath):
|
||||
actual = os.path.getsize(outpath)
|
||||
if expected == 0 or actual == expected:
|
||||
success = True; break
|
||||
# Size mismatch — might be in cache
|
||||
# If we got here, file downloaded but not at expected path
|
||||
# ModelScope puts it in local_dir; check again
|
||||
if os.path.exists(outpath):
|
||||
success = True; break
|
||||
# Retry with curl fallback
|
||||
if use_ms:
|
||||
base = f"https://modelscope.cn/api/v1/models/{REPO_MS}/repo?Revision=master&FilePath="
|
||||
else:
|
||||
base = f"https://huggingface.co/{REPO_HF}/resolve/main"
|
||||
if download_file_curl(fn, base, expected):
|
||||
success = True; break
|
||||
except Exception as e:
|
||||
if attempt < 2:
|
||||
print(f" {C.yel}retry {attempt+1}: {e}{C.r}")
|
||||
time.sleep(3)
|
||||
else:
|
||||
print(f" {C.yel}✗ failed: {e}{C.r}")
|
||||
|
||||
with lock:
|
||||
if success:
|
||||
completed.append(fn)
|
||||
elapsed = time.time() - t0
|
||||
have = sum(sizes.get(f,0) for f in completed)
|
||||
pct = 100.0 * have / total_bytes
|
||||
speed = (have - sum(sizes[f] for f in done_set)) / max(elapsed, 1)
|
||||
eta = (total_bytes - have) / speed if speed > 0 else 0
|
||||
print(f" {C.grn}✓{C.r} {fn} {C.dim}— {len(completed)}/{total} · "
|
||||
f"{pct:.1f}% · {fmt_time(elapsed)} · ETA {fmt_time(eta)}{C.r}")
|
||||
else:
|
||||
print(f" {C.yel}✗ GIVE UP: {fn}{C.r}")
|
||||
|
||||
threads = [threading.Thread(target=worker, args=(i,), daemon=True) for i in range(args.parallel)]
|
||||
for t in threads: t.start()
|
||||
for t in threads: t.join()
|
||||
|
||||
print()
|
||||
final = sum(1 for fn in shards
|
||||
if os.path.exists(os.path.join(DEST, fn))
|
||||
and os.path.getsize(os.path.join(DEST, fn)) == sizes.get(fn, 0))
|
||||
if final == total:
|
||||
print(f"{C.grn}{'='*50}")
|
||||
print(f" ✓ All {total} shards downloaded!{C.r}\n")
|
||||
else:
|
||||
print(f"{C.yel} {final}/{total} complete, {total-final} remaining{C.r}")
|
||||
print(f" Re-run to resume.\n")
|
||||
|
||||
if __name__ == "__main__":
|
||||
try: main()
|
||||
except KeyboardInterrupt:
|
||||
print(f"\n\n{C.yel}Interrupted. Re-run to resume — no data lost.{C.r}\n")
|
||||
@@ -1622,7 +1622,10 @@ static int expert_load(Model *m, int layer, int eid, ESlot *s, int fatal){
|
||||
* condvar exist ONLY to park/wake idle workers, never for correctness. Gated
|
||||
* behind PIPE=1; OFF => the original blocking-load + serial-matmul path runs
|
||||
* byte-identically. */
|
||||
static int g_pipe=0; /* PIPE=1: async expert-load pipeline (default OFF) */
|
||||
static int g_pipe=0; /* PIPE=1: async expert-load pipeline. Default ON for Windows
|
||||
* (parsed in main: getenv("PIPE")?:1 on _WIN32, :0 elsewhere).
|
||||
* Keeps expert pread off the forward-pass thread so loads overlap
|
||||
* the matmul. PIPE=0 opts back into the blocking serial path. */
|
||||
static int g_pipe_nw=8; /* PIPE_WORKERS=n: I/O worker threads (disk-parallel reads) */
|
||||
typedef struct {
|
||||
_Atomic uint64_t cur; /* (gen<<8)|index; gen main-only, index 0..njobs (≤64) */
|
||||
@@ -5079,7 +5082,13 @@ int main(int argc, char **argv){
|
||||
g_pilot_k = getenv("PILOT_K")?atoi(getenv("PILOT_K")):(g_pilot_real?6:8);
|
||||
if(g_pilot_k<1) g_pilot_k=1;
|
||||
g_disk_split = getenv("DISK_SPLIT")?atoi(getenv("DISK_SPLIT")):0; /* 1 = split dei disk load nelle stats */
|
||||
g_pipe = getenv("PIPE")?atoi(getenv("PIPE")):0; /* default OFF: overlap expert load ‖ matmul (byte-identical; reorders I/O). PIPE=1 opts in */
|
||||
g_pipe = getenv("PIPE")?atoi(getenv("PIPE")):
|
||||
#ifdef _WIN32
|
||||
1 /* default ON: overlap expert load ‖ matmul (byte-identical; reorders I/O). PIPE=0 opts out */
|
||||
#else
|
||||
0
|
||||
#endif
|
||||
;
|
||||
g_pipe_nw = getenv("PIPE_WORKERS")?atoi(getenv("PIPE_WORKERS")):8; /* I/O worker threads */
|
||||
if(g_pipe_nw<1) g_pipe_nw=1;
|
||||
g_direct = getenv("DIRECT")?atoi(getenv("DIRECT")):0;
|
||||
|
||||
@@ -51,6 +51,23 @@ int main(void){
|
||||
if(compat_open_direct("no_such_file.tmp")>=0) return fail("open missing file must fail");
|
||||
if(compat_fsize(-1)>=0) return fail("compat_fsize on bad fd must be negative");
|
||||
|
||||
/* compat_fadvise: WILLNEED warms the page cache (background read into throwaway
|
||||
* buffer), DONTNEED is a documented no-op. After a WILLNEED the buffered fd's
|
||||
* subsequent pread must still return the exact bytes — the cache-warmer must not
|
||||
* corrupt data. Bad fd / non-WILLNEED advice must be safe no-ops (return 0). */
|
||||
int wfd = open(TMPF, COMPAT_O_RDONLY);
|
||||
if(wfd<0) return fail("open buffered for fadvise");
|
||||
if(posix_fadvise(wfd, 0, (off_t)FSZ, POSIX_FADV_WILLNEED)!=0) return fail("WILLNEED returned nonzero");
|
||||
if(posix_fadvise(wfd, 0, (off_t)FSZ, POSIX_FADV_DONTNEED)!=0) return fail("DONTNEED should be a safe no-op (return 0)");
|
||||
if(posix_fadvise(-1, 0, (off_t)FSZ, POSIX_FADV_WILLNEED)!=0) return fail("WILLNEED on bad fd should no-op (return 0)");
|
||||
if(posix_fadvise(wfd, 0, 0, POSIX_FADV_WILLNEED)!=0) return fail("WILLNEED with len<=0 should no-op");
|
||||
/* verify data integrity through the buffered fd after the cache-warmer ran */
|
||||
uint8_t *verify=malloc(FSZ);
|
||||
if(pread(wfd, verify, FSZ, 0)!=(ssize_t)FSZ) return fail("fadvise: pread size");
|
||||
if(memcmp(verify, pat, FSZ)!=0) return fail("fadvise: data corrupted by cache-warmer");
|
||||
free(verify);
|
||||
close(wfd);
|
||||
|
||||
close(dfd);
|
||||
compat_aligned_free(buf); free(pat); remove(TMPF);
|
||||
puts("compat direct tests: ok");
|
||||
|
||||
@@ -116,7 +116,7 @@ def layer_idx(name):
|
||||
|
||||
def classify(name, n_layers, keep_mtp=False, keep_idx=False):
|
||||
if name.endswith("_scale_inv"): return "consumed" # FP8 base: gestito col suo peso
|
||||
# NVFP4 (modelopt): i sidecar delle scale sono consumati insieme al loro .weight U8.
|
||||
# NVFP4 (modelopt): i sidecar delle scale sono consumati insieme al loro U8 .weight.
|
||||
# EN: NVFP4 (modelopt): scale sidecars are consumed together with their U8 .weight.
|
||||
if name.endswith((".weight_scale", ".weight_scale_2", ".input_scale")): return "consumed"
|
||||
li = layer_idx(name)
|
||||
@@ -137,7 +137,20 @@ def classify(name, n_layers, keep_mtp=False, keep_idx=False):
|
||||
if name.endswith("norm.weight") or name == "model.norm.weight": return "f32"
|
||||
if name in ("model.embed_tokens.weight", "lm_head.weight"): return "io"
|
||||
if ".mlp.experts." in name and name.endswith(".weight"): return "x" # expert ROUTED (streaming)
|
||||
if name.endswith(".weight"): return "q" # attn/dense-mlp/shared (residente)
|
||||
# Split resident weights by type for mixed-precision control:
|
||||
# "sh" = shared expert (fires on every token, highest sensitivity)
|
||||
# "o" = o_proj attention (reconstructs output, biggest attn tensor)
|
||||
# "kvb" = kv_b_proj (reconstructs KV cache on every decode step)
|
||||
# "attn" = other attention projections (q_a, q_b, kv_a)
|
||||
# "dmlp" = dense MLP (first 3 layers)
|
||||
if "shared_experts" in name: return "sh"
|
||||
if name.endswith("o_proj.weight"): return "o"
|
||||
if name.endswith("kv_b_proj.weight"): return "kvb"
|
||||
if any(name.endswith(k) for k in ("q_a_proj.weight", "q_b_proj.weight",
|
||||
"kv_a_proj_with_mqa.weight")): return "attn"
|
||||
if any(name.endswith(k) for k in ("mlp.gate_proj.weight", "mlp.up_proj.weight",
|
||||
"mlp.down_proj.weight")): return "dmlp"
|
||||
if name.endswith(".weight"): return "q" # fallback: other resident weights
|
||||
return "f32"
|
||||
|
||||
# ---------- dequant NVFP4 (modelopt) di UN tensore expert -> f32 [O,I] ----------
|
||||
@@ -202,7 +215,7 @@ def dequant(f, name, keys):
|
||||
return f.get_tensor(name).to(torch.float32).numpy()
|
||||
|
||||
def convert_shard(path, out_dict, n_layers, ebits, io_bits, xbits,
|
||||
keep_mtp=False, keep_idx=False, group_size=0):
|
||||
keep_mtp=False, keep_idx=False, group_size=0, bits_map=None):
|
||||
from safetensors import safe_open
|
||||
with safe_open(path, framework="pt") as f:
|
||||
keys = set(f.keys())
|
||||
@@ -213,7 +226,15 @@ def convert_shard(path, out_dict, n_layers, ebits, io_bits, xbits,
|
||||
if kind == "f32":
|
||||
out_dict[name] = w.astype(np.float32)
|
||||
else:
|
||||
bits = io_bits if kind == "io" else xbits if kind == "x" else ebits
|
||||
# Resolve bits for this tensor type: use bits_map override if provided,
|
||||
# otherwise fall back to the classic ebits/xbits/io_bits scheme.
|
||||
if bits_map and kind in bits_map:
|
||||
bits = bits_map[kind]
|
||||
else:
|
||||
bits = io_bits if kind == "io" else xbits if kind == "x" else ebits
|
||||
# Any unknown kind that fell through classify as "q"
|
||||
if bits_map and kind not in bits_map and kind not in ("io", "x", "sh", "o", "kvb", "attn", "dmlp"):
|
||||
bits = ebits
|
||||
if w.ndim != 2: # es. bias 1D non previsto come 'q' -> tienilo f32
|
||||
out_dict[name] = w.astype(np.float32); continue
|
||||
if group_size > 0 and bits <= 4:
|
||||
@@ -234,6 +255,18 @@ def main():
|
||||
ap.add_argument("--ebits", type=int, default=None) # bit residenti (default 4; 8 per --mtp/--indexer)
|
||||
ap.add_argument("--io-bits", type=int, default=8) # bit di embed/lm_head
|
||||
ap.add_argument("--xbits", type=int, default=None) # bit degli expert ROUTED (streaming); default=ebits
|
||||
# Mixed-precision: per-tensor-type bit overrides. Default = ebits (all same).
|
||||
# Set these higher to protect sensitive tensors from quantization error.
|
||||
ap.add_argument("--shared-bits", type=int, default=None,
|
||||
help="bits for shared expert (fires on every token, highest sensitivity). Default=ebits")
|
||||
ap.add_argument("--o-bits", type=int, default=None,
|
||||
help="bits for o_proj attention (reconstructs output, biggest attn tensor). Default=ebits")
|
||||
ap.add_argument("--kvb-bits", type=int, default=None,
|
||||
help="bits for kv_b_proj (reconstructs KV cache on every decode). Default=ebits")
|
||||
ap.add_argument("--attn-bits", type=int, default=None,
|
||||
help="bits for other attention projections (q_a, q_b, kv_a). Default=ebits")
|
||||
ap.add_argument("--dmlp-bits", type=int, default=None,
|
||||
help="bits for dense MLP (first 3 layers). Default=ebits")
|
||||
ap.add_argument("--group-size", type=int, default=0, # 0 = per-row (backward compat); 128 = group-scaled
|
||||
help="group size for int4 scales: 0=per-row (default), 128=one scale per 128 elements (much better quality)")
|
||||
ap.add_argument("--n-layers", type=int, default=78)
|
||||
@@ -255,6 +288,17 @@ def main():
|
||||
a.ebits = 8 if (a.mtp or a.indexer) else 4
|
||||
if a.xbits is None: a.xbits = a.ebits
|
||||
|
||||
# Build per-type bits map. If a type-specific arg is set, use it; otherwise the
|
||||
# converter falls back to ebits for that type.
|
||||
bits_map = {}
|
||||
if a.shared_bits is not None: bits_map["sh"] = a.shared_bits
|
||||
if a.o_bits is not None: bits_map["o"] = a.o_bits
|
||||
if a.kvb_bits is not None: bits_map["kvb"] = a.kvb_bits
|
||||
if a.attn_bits is not None: bits_map["attn"] = a.attn_bits
|
||||
if a.dmlp_bits is not None: bits_map["dmlp"] = a.dmlp_bits
|
||||
if bits_map:
|
||||
print(f"[MIXED] precision map: " + ", ".join(f"{k}={v}bit" for k,v in sorted(bits_map.items())))
|
||||
|
||||
if a.selftest_nvfp4:
|
||||
import torch
|
||||
# 1) LUT e2m1: i 16 codici devono decodificare esattamente ai valori attesi.
|
||||
@@ -336,7 +380,7 @@ def main():
|
||||
shards = sorted(glob.glob(os.path.join(a.indir, "*.safetensors")))
|
||||
from safetensors.numpy import save_file
|
||||
for i, sp in enumerate(shards):
|
||||
out = {}; convert_shard(sp, out, a.n_layers, a.ebits, a.io_bits, a.xbits, group_size=a.group_size)
|
||||
out = {}; convert_shard(sp, out, a.n_layers, a.ebits, a.io_bits, a.xbits, group_size=a.group_size, bits_map=bits_map)
|
||||
save_file(out, os.path.join(a.outdir, f"out-{i:05d}.safetensors"))
|
||||
# copia config + tokenizer
|
||||
for fn in ["config.json"]:
|
||||
@@ -579,7 +623,7 @@ def main():
|
||||
if os.path.exists(outp): print(f"[MTP] {outp} already done"); continue
|
||||
print(f"[MTP {i+1}/{len(mtp_shards)}] downloading {sh}...", flush=True)
|
||||
p = download_retry(a.repo, sh, tmp)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, keep_mtp=True, group_size=a.group_size)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, keep_mtp=True, group_size=a.group_size, bits_map=bits_map)
|
||||
save_file(out, outp)
|
||||
os.remove(p)
|
||||
for blob in glob.glob(os.path.join(tmp, "**", "*"), recursive=True):
|
||||
@@ -599,7 +643,7 @@ def main():
|
||||
if os.path.exists(outp): continue # gia' fatto -> ripartibile
|
||||
print(f"[IDX {i+1}/{len(idx_shards)}] downloading {sh}...", flush=True)
|
||||
p = download_retry(a.repo, sh, tmp)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, keep_idx=True, group_size=a.group_size)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, keep_idx=True, group_size=a.group_size, bits_map=bits_map)
|
||||
if out: save_file(out, outp)
|
||||
os.remove(p)
|
||||
for blob in glob.glob(os.path.join(tmp, "**", "*"), recursive=True):
|
||||
@@ -613,7 +657,7 @@ def main():
|
||||
if os.path.exists(outp): continue # gia' fatto -> ripartibile
|
||||
print(f"[{i+1}/{len(shards)}] downloading {sh} ({free_gb(a.outdir):.0f} GB free)...", flush=True)
|
||||
p = download_retry(a.repo, sh, tmp)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, group_size=a.group_size)
|
||||
out = {}; convert_shard(p, out, a.n_layers, a.ebits, a.io_bits, a.xbits, group_size=a.group_size, bits_map=bits_map)
|
||||
save_file(out, outp)
|
||||
os.remove(p) # <-- cancella subito lo shard fp8
|
||||
for blob in glob.glob(os.path.join(tmp, "**", "*"), recursive=True):
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Helper: salva pesi in FP8 e4m3 + scale a blocchi 128x128, nello STESSO layout del
|
||||
checkpoint reale GLM-5.2-FP8 che `convert_fp8_to_int4.py` legge.
|
||||
|
||||
Layout (deve combaciare col `dequant()` del converter, convert_fp8_to_int4.py:164-169):
|
||||
- `name` F8_E4M3 [O, I]
|
||||
- `name_scale_inv` F32 [ceil(O/128), ceil(I/128)] (NOTA: '_scale_inv', underscore)
|
||||
dequant: W = q.float() * scale.repeat_interleave(128,0).repeat_interleave(128,1)[:O,:I]
|
||||
|
||||
Convenzione FBGEMM/TransformerEngine: scale = amax(blocco)/448 (448 = max e4m3),
|
||||
si MEMORIZZA il valore e si MOLTIPLICA in dequant. Malgrado il nome "_scale_inv" il
|
||||
checkpoint memorizza la scala (non il reciproco): e' un MOLTIPLIER.
|
||||
|
||||
EN: Helper that writes weights as FP8 e4m3 with 128x128 block scales, in the SAME layout
|
||||
EN: as the real GLM-5.2-FP8 checkpoint that `convert_fp8_to_int4.py` reads.
|
||||
EN: FBGEMM/TransformerEngine convention: scale = amax(block)/448, stored (not its
|
||||
EN: reciprocal) and MULTIPLIED on dequant. Despite the name "_scale_inv" it is a multiplier.
|
||||
"""
|
||||
import torch
|
||||
|
||||
E4M3_MAX = 448.0 # max valore rappresentabile in float8_e4m3fn / max representable value
|
||||
BLOCK = 128 # granularita' delle scale a blocchi del checkpoint FP8 / FP8 block scale granularity
|
||||
|
||||
|
||||
def keep_f32(name, t):
|
||||
"""Stesso set F32 di `classify()` in convert_fp8_to_int4.py (norme, router, bias 1-D).
|
||||
Tutti gli altri tensori 2-D vengono quantizzati FP8 (attn/mlp/shared/expert/embed/lm_head).
|
||||
EN: Same F32 set as the converter's classify(): norms, router, 1-D biases. All other 2-D
|
||||
EN: tensors are FP8-quantized (attn/mlp/shared/expert/embed/lm_head)."""
|
||||
if t.dim() < 2:
|
||||
return True # bias 1-D, e_score_correction_bias
|
||||
if name.endswith("e_score_correction_bias"):
|
||||
return True
|
||||
if name.endswith("mlp.gate.weight"):
|
||||
return True # router (NON gate_proj): tenuto F32 / kept F32
|
||||
if name.endswith("norm.weight") or name == "model.norm.weight":
|
||||
return True # RMSNorm
|
||||
return False
|
||||
|
||||
|
||||
def fp8_block_quantize(w):
|
||||
"""w: [O,I] f32 -> (w_fp8 float8_e4m3fn [O,I], scale_inv f32 [ceil(O/128),ceil(I/128)]).
|
||||
Identica matematica al `--selftest` del converter (scale = amax(blocco)/448). Padda a
|
||||
multipli di 128 internamente (gli zeri non alzano l'amax) e fa slice al risultato.
|
||||
EN: same math as the converter's --selftest. Pads to 128 multiples internally (zeros do
|
||||
EN: not raise amax), slices the result back to [O,I]."""
|
||||
O, I = w.shape
|
||||
nbO, nbI = (O + BLOCK - 1) // BLOCK, (I + BLOCK - 1) // BLOCK
|
||||
Op, Ip = nbO * BLOCK, nbI * BLOCK
|
||||
wpad = torch.zeros(Op, Ip, dtype=torch.float32, device=w.device)
|
||||
wpad[:O, :I] = w
|
||||
wb = wpad.view(nbO, BLOCK, nbI, BLOCK) # [nbO, BLOCK, nbI, BLOCK]
|
||||
amax = wb.abs().amax(dim=(1, 3)) # [nbO, nbI]
|
||||
scale = amax / E4M3_MAX # FBGEMM/TE: memorizza la scala / store the scale
|
||||
scale = torch.where(scale == 0, torch.ones_like(scale), scale) # blocco tutto-zero -> no div0
|
||||
scale = scale.to(torch.float32)
|
||||
q = (wpad / scale.repeat_interleave(BLOCK, 0).repeat_interleave(BLOCK, 1)).clamp(-E4M3_MAX, E4M3_MAX)
|
||||
w_fp8 = q.to(torch.float8_e4m3fn)
|
||||
return w_fp8[:O, :I].contiguous(), scale.contiguous()
|
||||
|
||||
|
||||
def fp8_block_dequantize(w_fp8, scale):
|
||||
"""Esatto inverso di fp8_block_quantize, e identico al `dequant()` del converter.
|
||||
EN: exact inverse of fp8_block_quantize, identical to the converter's dequant()."""
|
||||
O, I = w_fp8.shape
|
||||
qf = w_fp8.to(torch.float32)
|
||||
return qf * scale.repeat_interleave(BLOCK, 0).repeat_interleave(BLOCK, 1)[:O, :I]
|
||||
|
||||
|
||||
def unfuse_experts(sd):
|
||||
"""Split HF's fused 3-D `experts.gate_up_proj` [E, 2*M, I] into per-expert 2-D
|
||||
`experts.{e}.gate_proj` [M, I] + `experts.{e}.up_proj` [M, I], and
|
||||
`experts.down_proj` [E, I, M] -> `experts.{e}.down_proj` [M_out, I].
|
||||
|
||||
The real GLM-5.2-FP8 checkpoint stores experts UNFUSED as per-expert 2-D tensors
|
||||
(gate_proj, up_proj, down_proj), each with its own _scale_inv. HF's
|
||||
GlmMoeDsaForCausalLM fuses gate+up into a single 3-D gate_up_proj for efficiency.
|
||||
The converter (classify + ndim!=2 guard) and the C engine both expect the unfused
|
||||
layout, so we split before saving.
|
||||
|
||||
Idempotent: if experts are already unfused (no 3-D gate_up_proj), returns sd as-is.
|
||||
EN: split HF's fused 3-D expert weights into the per-expert 2-D layout that the real
|
||||
EN: checkpoint uses and the converter/engine expect. No-op if already unfused."""
|
||||
keys_to_remove = []
|
||||
new_entries = {}
|
||||
for name, t in sd.items():
|
||||
if not name.endswith(".mlp.experts.gate_up_proj"):
|
||||
continue
|
||||
# prefix = everything before ".mlp.experts.gate_up_proj"
|
||||
prefix = name[:-len(".mlp.experts.gate_up_proj")]
|
||||
E, twoM, I = t.shape # [E, 2*intermediate, input]
|
||||
M = twoM // 2
|
||||
for e in range(E):
|
||||
new_entries[f"{prefix}.mlp.experts.{e}.gate_proj.weight"] = t[e, :M, :].contiguous()
|
||||
new_entries[f"{prefix}.mlp.experts.{e}.up_proj.weight"] = t[e, M:, :].contiguous()
|
||||
keys_to_remove.append(name)
|
||||
# down_proj may be 3-D [E, I, M] in the fused form, or already per-expert
|
||||
for name, t in sd.items():
|
||||
if not name.endswith(".mlp.experts.down_proj") or t.dim() != 3:
|
||||
continue
|
||||
prefix = name[:-len(".mlp.experts.down_proj")]
|
||||
E = t.shape[0]
|
||||
for e in range(E):
|
||||
new_entries[f"{prefix}.mlp.experts.{e}.down_proj.weight"] = t[e].contiguous()
|
||||
keys_to_remove.append(name)
|
||||
for k in keys_to_remove:
|
||||
sd.pop(k, None)
|
||||
sd.update(new_entries)
|
||||
return sd
|
||||
|
||||
|
||||
def state_dict_to_fp8(sd):
|
||||
"""Converte uno state_dict HuggingFace nel layout FP8 del checkpoint reale:
|
||||
per ogni tensore quantizzabile 2-D scrive `{name}` (F8_E4M3) + `{name}_scale_inv` (F32);
|
||||
norme/router/bias e qualsiasi tensore NON 2-D (es. pesi MLA impaccati 3-D) restano nel
|
||||
dtype originale. Questo rispecchia il guard `w.ndim != 2 -> f32` del converter
|
||||
(convert_fp8_to_int4.py:184). EN: builds the real-checkpoint FP8 layout. Only exactly-2-D
|
||||
tensors are FP8-quantized; anything else (1-D, 3-D packed MLA weights, ...) is kept, exactly
|
||||
like the converter's `ndim != 2 -> f32` guard."""
|
||||
out = {}
|
||||
for name, t in sd.items():
|
||||
if keep_f32(name, t) or t.dim() != 2:
|
||||
out[name] = t # f32 / 1-D / 3-D+: tieni / keep
|
||||
else:
|
||||
w_fp8, scale = fp8_block_quantize(t.float())
|
||||
out[name] = w_fp8
|
||||
out[name + "_scale_inv"] = scale
|
||||
return out
|
||||
|
||||
|
||||
def save_fp8_safetensors(sd, path):
|
||||
"""Quantizza a blocchi FP8 e salva in un singolo safetensors leggibile dal converter
|
||||
via `--indir`. EN: block-quantize to FP8 and save a single safetensors for the converter."""
|
||||
from safetensors.torch import save_file
|
||||
out = state_dict_to_fp8(sd)
|
||||
save_file({k: v.contiguous() for k, v in out.items()}, str(path))
|
||||
n_fp8 = sum(1 for v in out.values() if v.dtype == torch.float8_e4m3fn)
|
||||
return n_fp8, len(out)
|
||||
@@ -3,15 +3,27 @@
|
||||
This is not a useful language model. It preserves the real glm_moe_dsa data
|
||||
flow while remaining small enough to generate locally and run repeated CPU/CUDA
|
||||
A/B tests without downloading the 379 GB checkpoint.
|
||||
|
||||
With --fp8 the weights are written as FP8 e4m3 + 128x128 block scale_inv, in the
|
||||
SAME layout as the real GLM-5.2-FP8 checkpoint, so convert_fp8_to_int4.py can
|
||||
exercise its FP8->int4 dequant path on a local fixture (its dims are 128-friendly,
|
||||
so this is also the right fixture for --group-size 128 testing):
|
||||
|
||||
python tools/make_glm_bench_model.py --fp8 --output glm_bench_fp8
|
||||
python tools/convert_fp8_to_int4.py --indir glm_bench_fp8 --outdir glm_bench_i4 --ebits 4 --group-size 128
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from transformers import GlmMoeDsaConfig, GlmMoeDsaForCausalLM
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent)) # importa glm_fp8_emit se lanciato da c/
|
||||
from glm_fp8_emit import save_fp8_safetensors, unfuse_experts
|
||||
|
||||
|
||||
def build_config() -> GlmMoeDsaConfig:
|
||||
return GlmMoeDsaConfig(
|
||||
@@ -51,6 +63,9 @@ def main() -> None:
|
||||
parser.add_argument("--output", default="glm_bench_medium")
|
||||
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
||||
parser.add_argument("--seed", type=int, default=1234)
|
||||
parser.add_argument("--fp8", action="store_true",
|
||||
help="write weights as FP8 e4m3 + 128x128 block scale_inv (same layout as "
|
||||
"GLM-5.2-FP8) instead of bf16, so convert_fp8_to_int4.py can dequant+requant")
|
||||
args = parser.parse_args()
|
||||
|
||||
torch.manual_seed(args.seed)
|
||||
@@ -70,7 +85,6 @@ def main() -> None:
|
||||
output = Path(args.output)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
params = sum(p.numel() for p in model.parameters())
|
||||
model.save_pretrained(output, safe_serialization=True, max_shard_size="4GB")
|
||||
|
||||
model.to(args.device)
|
||||
prompt = [3, 14, 159, 26, 53, 58, 200, 11, 77, 240, 5, 99]
|
||||
@@ -79,6 +93,25 @@ def main() -> None:
|
||||
full = model.generate(ids, max_new_tokens=8, do_sample=False, use_cache=True)[0]
|
||||
logits = model(full.unsqueeze(0), use_cache=False).logits[0]
|
||||
|
||||
# Unfuse experts AFTER reference generation (model needs fused weights for
|
||||
# forward/generate) but BEFORE saving — the real checkpoint and the converter
|
||||
# + C engine all expect per-expert 2-D gate_proj/up_proj/down_proj tensors.
|
||||
sd = model.state_dict()
|
||||
unfuse_experts(sd)
|
||||
|
||||
if args.fp8:
|
||||
n_fp8, n_tot = save_fp8_safetensors(sd, output / "model.safetensors")
|
||||
# save_pretrained scrive config.json; nel path FP8 lo bypassiamo, quindi lo scriviamo
|
||||
# a mano (serve al converter e al motore C). EN: save_pretrained writes config.json;
|
||||
# the FP8 path bypasses it, so write it manually (converter + C engine need it).
|
||||
(output / "config.json").write_text(json.dumps(cfg.to_dict()))
|
||||
print(f"saved FP8: {n_fp8} e4m3 tensors (+{n_tot - n_fp8} scale_inv sidecars / f32) "
|
||||
f"-> {output / 'model.safetensors'}")
|
||||
else:
|
||||
from safetensors.torch import save_file
|
||||
save_file({k: v.contiguous() for k, v in sd.items()}, str(output / "model.safetensors"))
|
||||
(output / "config.json").write_text(json.dumps(cfg.to_dict()))
|
||||
|
||||
ref = {
|
||||
"prompt_ids": prompt,
|
||||
"full_ids": full.cpu().tolist(),
|
||||
@@ -89,6 +122,7 @@ def main() -> None:
|
||||
"seed": args.seed,
|
||||
"parameters": params,
|
||||
"parameters_billions": round(params / 1e9, 4),
|
||||
"format": "fp8-e4m3-128" if args.fp8 else "bf16",
|
||||
"purpose": "backend benchmark fixture; random weights, not a language model",
|
||||
}
|
||||
(output / "bench_manifest.json").write_text(json.dumps(manifest, indent=2))
|
||||
|
||||
@@ -3,10 +3,34 @@ Architettura vera (MLA + DSA indexer + router sigmoid/noaux_tc + shared expert),
|
||||
dimensioni minuscole. Salva pesi+config in c/glm_tiny/ e un riferimento greedy in
|
||||
c/ref_glm.json. seq corta (<= index_topk) cosi' il DSA seleziona tutte le key e
|
||||
l'attenzione coincide con la MLA densa: il motore C puo' validare senza implementare
|
||||
l'indexer sparso."""
|
||||
import json, torch
|
||||
l'indexer sparso.
|
||||
|
||||
--fp8: salva i pesi come FP8 e4m3 + scale a blocchi 128x128 (layout del checkpoint reale
|
||||
GLM-5.2-FP8) invece di bf16, cosi' convert_fp8_to_int4.py puo' esercitare il path FP8->int4
|
||||
su un modello minuscolo. PRIMA di calcolare ref_glm.json fa il round-trip dei pesi per FP8
|
||||
(quant->dequant, copy_ nel modello): cosi' il riferimento riflette ESATTAMENTE il modello
|
||||
FP8 che il converter legge, non il modello bf16 a precisione piena. Default: bf16 (oracolo
|
||||
originale invariato).
|
||||
EN: --fp8 writes FP8 e4m3 + 128x128 block scale_inv (real GLM-5.2-FP8 layout) instead of bf16,
|
||||
EN: so convert_fp8_to_int4.py can run its FP8->int4 path on a tiny model. ref_glm.json is
|
||||
EN: computed AFTER the FP8 round-trip, so the reference matches exactly what the converter
|
||||
EN: ingests. Default: bf16 (original oracle unchanged)."""
|
||||
import json, sys, argparse
|
||||
from pathlib import Path
|
||||
import torch
|
||||
from transformers import GlmMoeDsaConfig, GlmMoeDsaForCausalLM
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent)) # importa glm_fp8_emit se lanciato da c/
|
||||
from glm_fp8_emit import (fp8_block_quantize, fp8_block_dequantize, keep_f32,
|
||||
save_fp8_safetensors, unfuse_experts)
|
||||
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--fp8", action="store_true",
|
||||
help="salva in FP8 e4m3 + 128x128 block scale_inv (layout GLM-5.2-FP8) e "
|
||||
"calcola ref_glm.json sul modello dopo il round-trip FP8. "
|
||||
"EN: write FP8 e4m3 + block scale_inv, ref computed on FP8-rounded model")
|
||||
args = ap.parse_args()
|
||||
|
||||
torch.manual_seed(1234)
|
||||
|
||||
cfg = GlmMoeDsaConfig(
|
||||
@@ -53,6 +77,18 @@ with torch.no_grad():
|
||||
layer.mlp.gate.e_score_correction_bias.copy_(
|
||||
torch.linspace(-0.1, 0.1, cfg.n_routed_experts))
|
||||
|
||||
# --fp8: round-trip dei pesi quantizzabili per FP8 PRIMA di calcolare il riferimento,
|
||||
# cosi' ref_glm.json riflette esattamente il modello FP8 che il converter leggera'.
|
||||
# Norme/router/bias (keep_f32) restano a precisione piena. EN: --fp8: round-trip quantizable
|
||||
# weights through FP8 before computing the reference, so ref_glm.json matches the FP8 model.
|
||||
if args.fp8:
|
||||
with torch.no_grad():
|
||||
for n, p in model.named_parameters():
|
||||
if keep_f32(n, p) or p.dim() != 2:
|
||||
continue
|
||||
q, s = fp8_block_quantize(p)
|
||||
p.copy_(fp8_block_dequantize(q, s))
|
||||
|
||||
print("=== state_dict tensors (names used by the C loader) ===")
|
||||
for n, p in model.state_dict().items():
|
||||
print(f" {n:60s} {tuple(p.shape)}")
|
||||
@@ -73,7 +109,20 @@ with torch.no_grad():
|
||||
tf_pred = lg.argmax(-1).tolist()
|
||||
print("tf_pred:", tf_pred)
|
||||
|
||||
model.save_pretrained("glm_tiny", safe_serialization=True)
|
||||
# Unfuse experts AFTER reference generation (model needs fused weights for
|
||||
# forward/generate) but BEFORE saving — the real checkpoint and the converter
|
||||
# + C engine all expect per-expert 2-D gate_proj/up_proj/down_proj tensors.
|
||||
sd = model.state_dict()
|
||||
unfuse_experts(sd)
|
||||
|
||||
if args.fp8:
|
||||
n_fp8, n_tot = save_fp8_safetensors(sd, "glm_tiny/model.safetensors")
|
||||
print(f"\nsaved FP8: {n_fp8} e4m3 tensors (+{n_tot - n_fp8} scale_inv sidecars / f32) "
|
||||
f"-> glm_tiny/model.safetensors")
|
||||
else:
|
||||
from safetensors.torch import save_file
|
||||
save_file({k: v.contiguous() for k, v in sd.items()}, "glm_tiny/model.safetensors")
|
||||
json.dump(cfg.to_dict(), open("glm_tiny/config.json", "w"))
|
||||
json.dump({"prompt_ids": prompt, "full_ids": full, "tf_pred": tf_pred}, open("ref_glm.json", "w"))
|
||||
print("\nsaved: glm_tiny/ (weights + config) and ref_glm.json")
|
||||
print("saved: glm_tiny/ (weights + config) and ref_glm.json"
|
||||
+ (" [fp8]" if args.fp8 else ""))
|
||||
|
||||
Reference in New Issue
Block a user