kingjones777's picture
Add files using upload-large-folder tool
a9b5915 verified
Raw History Blame
6.16 kB
#!/usr/bin/env python3
"""Agnes-3.0-Flash Preview -> stock Qwen3.5 HF checkpoint for convert_hf_to_gguf.py.
Every transform is exact (verified separately by verify_fold.py):
1. *.delta_attn.* -> *.linear_attn.* the converter V-head-reorders ONLY linear_attn.* names
2. *.global_attn.* -> *.self_attn.*
3. mlp.parallel_ffn folded into mlp reference: y = down(act(gate x)*up x) + parallel_ffn(x)
gate/up: cat(main, par, dim=0) plain sum, same activation => concatenation is exact
down : cat(main, par, dim=1)
4. mtp.layers.N.mlp zero-padded to the folded width zero SwiGLU rows contribute exactly 0
5. config: qwen3_5 model/arch ids, layer_types mapped, full_attention_interval,
intermediate_size = main + parallel
usage: fold_agnes.py <agnes_hf_dir> <out_dir>
"""
import hashlib, json, os, re, shutil, sys, time
import torch
from safetensors import safe_open
from safetensors.torch import save_file
SRC, DST = sys.argv[1], sys.argv[2]
os.makedirs(DST, exist_ok=True)
cfg = json.load(open(os.path.join(SRC, "config.json")))
tc = cfg["text_config"]
MAIN = int(tc["intermediate_size"])
PAR = int(tc["parallel_ffn_intermediate_size"])
FOLD = MAIN + PAR
N_LAYERS = int(tc["num_hidden_layers"])
INTERVAL = int(tc["global_attention_interval"])
LT = {"agnes_delta_attention": "linear_attention", "agnes_global_attention": "full_attention"}
# layer plan must be the pattern the loader will re-derive from full_attention_interval
plan = tc["layer_types"]
assert len(plan) == N_LAYERS, (len(plan), N_LAYERS)
for i, t in enumerate(plan):
want = "agnes_global_attention" if (i + 1) % INTERVAL == 0 else "agnes_delta_attention"
assert t == want, f"layer {i}: {t} != {want} (loader derives the pattern from the interval)"
wm = json.load(open(os.path.join(SRC, "model.safetensors.index.json")))["weight_map"]
par_names = {k for k in wm if ".mlp.parallel_ffn." in k}
par_files = sorted({wm[k] for k in par_names})
main_files = sorted({f for f in wm.values()} - set(par_files))
assert len(par_names) == 3 * N_LAYERS, len(par_names)
assert not any(".mlp.parallel_ffn." not in k for k in wm if wm[k] in par_files), "parallel file holds other tensors"
par_handles = {f: safe_open(os.path.join(SRC, f), "pt") for f in par_files}
def par(name):
return par_handles[wm[name]].get_tensor(name)
RE_MAIN_MLP = re.compile(r"^(model\.language_model\.layers\.(\d+)\.mlp)\.(gate_proj|up_proj|down_proj)\.weight$")
RE_MTP_MLP = re.compile(r"^mtp\.layers\.\d+\.mlp\.(gate_proj|up_proj|down_proj)\.weight$")
def rename(n):
n = n.replace(".delta_attn.", ".linear_attn.")
n = n.replace(".global_attn.", ".self_attn.")
return n
new_map, total, stats = {}, 0, {"fold": 0, "pad": 0, "pass": 0}
t0 = time.time()
for f in main_files:
out = {}
with safe_open(os.path.join(SRC, f), "pt") as h:
for name in h.keys():
t = h.get_tensor(name)
m = RE_MAIN_MLP.match(name)
if m:
prefix, _, kind = m.groups()
p = par(f"{prefix}.parallel_ffn.{kind}.weight")
assert t.dtype == p.dtype, (name, t.dtype, p.dtype)
if kind in ("gate_proj", "up_proj"):
assert tuple(t.shape) == (MAIN, t.shape[1]) and tuple(p.shape) == (PAR, t.shape[1]), (name, t.shape, p.shape)
t = torch.cat([t, p], dim=0)
else:
assert tuple(t.shape) == (t.shape[0], MAIN) and tuple(p.shape) == (t.shape[0], PAR), (name, t.shape, p.shape)
t = torch.cat([t, p], dim=1)
stats["fold"] += 1
elif RE_MTP_MLP.match(name):
kind = RE_MTP_MLP.match(name).group(1)
if kind in ("gate_proj", "up_proj"):
assert t.shape[0] == MAIN, (name, t.shape)
t = torch.cat([t, torch.zeros((PAR, t.shape[1]), dtype=t.dtype)], dim=0)
else:
assert t.shape[1] == MAIN, (name, t.shape)
t = torch.cat([t, torch.zeros((t.shape[0], PAR), dtype=t.dtype)], dim=1)
stats["pad"] += 1
else:
stats["pass"] += 1
nn_ = rename(name)
assert nn_ not in out and nn_ not in new_map, f"name collision {nn_}"
out[nn_] = t.contiguous()
save_file(out, os.path.join(DST, f), metadata={"format": "pt"})
for k, v in out.items():
new_map[k] = f
total += v.numel() * v.element_size()
print(f" wrote {f}: {len(out)} tensors ({time.time()-t0:.0f}s)", flush=True)
del out
json.dump({"metadata": {"total_size": total}, "weight_map": dict(sorted(new_map.items()))},
open(os.path.join(DST, "model.safetensors.index.json"), "w"), indent=2)
# ---- config ----
o = json.loads(json.dumps(cfg))
o.pop("auto_map", None)
o["architectures"] = ["Qwen3_5ForConditionalGeneration"]
o["model_type"] = "qwen3_5"
t = o["text_config"]
t["model_type"] = "qwen3_5_text"
t["layer_types"] = [LT[x] for x in t["layer_types"]]
t["full_attention_interval"] = t.pop("global_attention_interval")
t["intermediate_size"] = FOLD
t.pop("parallel_ffn_intermediate_size")
o["vision_config"]["model_type"] = "qwen3_5"
json.dump(o, open(os.path.join(DST, "config.json"), "w"), indent=2)
# ---- sidecar files the converter/tokenizer read (custom code + sglang patch deliberately NOT copied) ----
KEEP = ["tokenizer.json", "tokenizer_config.json", "vocab.json", "merges.txt", "chat_template.jinja",
"generation_config.json", "preprocessor_config.json", "video_preprocessor_config.json",
"LICENSE", "README.md"]
for k in KEEP:
if os.path.exists(os.path.join(SRC, k)):
shutil.copy2(os.path.join(SRC, k), os.path.join(DST, k))
prov = {"source_repo": "Agnes-AI/Agnes-3.0-Flash", "source_dir": SRC, "main_ffn": MAIN, "parallel_ffn": PAR,
"folded_ffn": FOLD, "stats": stats, "tensors_out": len(new_map), "bytes_out": total,
"fold_script_sha256": hashlib.sha256(open(__file__, "rb").read()).hexdigest()}
json.dump(prov, open(os.path.join(DST, "FOLD_PROVENANCE.json"), "w"), indent=2)
print("DONE", json.dumps(prov), flush=True)