File size: 6,156 Bytes
a9b5915
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
#!/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)