File size: 13,189 Bytes
00bd945
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fd26a7a
fc2119a
fd26a7a
 
 
fc2119a
00bd945
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aec16a2
 
 
 
 
 
 
 
 
 
 
 
 
00bd945
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aec16a2
 
00bd945
 
 
 
aec16a2
00bd945
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
"""Plain data-free NVFP4; deterministic source-supported FFN folding."""
import collections
import copy
import gc
import hashlib
import importlib.metadata
import json
import os
import re
import shutil
import time
from pathlib import Path

import torch
from huggingface_hub import HfApi, snapshot_download
from safetensors import safe_open
from safetensors.torch import save_file
from transformers import AutoModelForImageTextToText
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.nn import TensorQuantizer
from modelopt.torch.export import export_hf_checkpoint

SOURCE = 'Agnes-AI/Agnes-3.0-Flash'
REV = '24f712ce59379b54c4a141d2708c35daf5ff613b'
TARGET = 'ProCreations/Agnes-3.0-Flash-NVFP4'
ROOT = Path('/workspace/agnes')
OUTPUT = ROOT/'export'
API = HfApi()
SELECT = re.compile(r'^model\.language_model\.layers\.\d+\.(?:mlp\.(?:gate_proj|up_proj|down_proj)|global_attn\.(?:q_proj|k_proj|v_proj|o_proj))\.weight$')
REPORT = dict(source=SOURCE, source_revision=REV, target=TARGET,
    status='building', code_revision=os.environ['AGNES_CODE_REVISION'],
    modelopt_revision='5cae3940402f1ced98069a666b0bec72ec8b33b5',
    quantization=dict(method='plain NVFP4 max', training=False, calibration_examples=0,
        calibration_tokens=0, calibration_forward_calls=0, activation_global_scale=1.0,
        block_size=16, activation_block_scales='dynamic E4M3', kv_cache_quantized=False),
    folding=dict(main_width=17408, parallel_width=2048, folded_width=19456,
        gate_up_dimension=0, down_dimension=1, layers=72,
        reason='Match the upstream SGLang BF16 loader before quantization'),
    evaluation=dict(status='pending'))

def record():
    (ROOT/'quality_report.json').write_text(json.dumps(REPORT, indent=2)+'\n')
    API.upload_file(repo_id=TARGET, path_or_fileobj=str(ROOT/'quality_report.json'),
        path_in_repo='quality_report.json', commit_message='Record plain NVFP4 build checks')

def selected(name):
    return bool(SELECT.fullmatch(name))

def main():
    assert API.model_info(TARGET).private
    assert not any(n.endswith('.safetensors') for n in API.list_repo_files(TARGET))
    ROOT.mkdir(parents=True, exist_ok=True)
    torch.set_num_threads(8)
    torch.manual_seed(20260912)
    REPORT['versions'] = {p:importlib.metadata.version(p) for p in ['torch','transformers','nvidia-modelopt','accelerate']}
    REPORT['hardware'] = [torch.cuda.get_device_name(0)]
    record()
    from transformers.dynamic_module_utils import get_class_from_dynamic_module
    source = Path(snapshot_download(SOURCE, revision=REV, local_dir=ROOT/'source', allow_patterns=['*.py','config.json'], max_workers=8))
    cls = get_class_from_dynamic_module('modeling_agnes.AgnesForConditionalGeneration',source)
    assert cls.__name__ == 'AgnesForConditionalGeneration'
    print('AGNES_REMOTE_IMPLEMENTATION_IMPORT_PASS',flush=True)
    source = Path(snapshot_download(SOURCE, revision=REV, local_dir=ROOT/'source', max_workers=8))
    src_index = json.loads((source/'model.safetensors.index.json').read_text())
    src_map = src_index['weight_map']
    model = AutoModelForImageTextToText.from_pretrained(source, trust_remote_code=True,
        dtype=torch.bfloat16, device_map={'':0}, attn_implementation='sdpa').eval()
    assert all(p.device.type == 'cuda' for p in model.parameters())
    assert len(model.model.language_model.layers) == 72
    with torch.no_grad():
        for layer in model.model.language_model.layers:
            mlp = layer.mlp
            assert mlp.parallel_ffn is not None
            for name in ['gate_proj','up_proj','down_proj']:
                original, branch = getattr(mlp,name), getattr(mlp.parallel_ffn,name)
                dim = 1 if name == 'down_proj' else 0
                merged = torch.cat([original.weight,branch.weight], dim=dim)
                # Check both slices before discarding the separate branches.
                a,b = merged.split([original.weight.shape[dim],branch.weight.shape[dim]],dim=dim)
                assert torch.equal(a,original.weight) and torch.equal(b,branch.weight)
                replacement = torch.nn.Linear(merged.shape[1],merged.shape[0],bias=False,
                    device=merged.device,dtype=merged.dtype)
                replacement.weight = torch.nn.Parameter(merged,requires_grad=False)
                setattr(mlp,name,replacement)
            mlp.parallel_ffn = None
    text = model.config.text_config
    text.agnes_original_intermediate_size = 17408
    text.agnes_original_parallel_ffn_intermediate_size = 2048
    text.intermediate_size = 19456
    text.parallel_ffn_intermediate_size = 0
    gc.collect(); torch.cuda.empty_cache()
    preset = copy.deepcopy(mtq.NVFP4_DEFAULT_CFG)
    weight = next(x['cfg'] for x in preset['quant_cfg'] if x.get('quantizer_name')=='*weight_quantizer')
    activation = next(x['cfg'] for x in preset['quant_cfg'] if x.get('quantizer_name')=='*input_quantizer')
    activation['constant_amax'] = 2688.0
    rules = [dict(quantizer_name='*',enable=False)]
    targets = [name for name,module in model.named_modules()
               if isinstance(module,torch.nn.Linear) and selected(name+'.weight')]
    assert len(targets) == 288, len(targets)
    for name in targets:
        rules += [dict(quantizer_name=name+'.weight_quantizer',cfg=weight),
                  dict(quantizer_name=name+'.input_quantizer',cfg=activation)]
    cfg = dict(quant_cfg=rules,algorithm=dict(method='max',layerwise=dict(enable=False),
        skip_forward_without_activation_calib=True))
    REPORT['quantization']['configuration'] = cfg
    def forbid(*args):
        REPORT['quantization']['calibration_forward_calls'] += 1
        raise RuntimeError('Data-free quantization may not execute calibration forwards')
    hook = model.register_forward_pre_hook(forbid)
    started = time.monotonic()
    with torch.inference_mode():
        # AgnesDeltaAttention ends in "Attention" and its Python module exposes
        # ALL_ATTENTION_FUNCTIONS, so ModelOpt's heuristic incorrectly wraps it
        # as a standard attention class. KV quantization is explicitly out of
        # scope: omit that optional registration callback for this conversion.
        from modelopt.torch.quantization.plugins.custom import CUSTOM_MODEL_PLUGINS
        from modelopt.torch.quantization.plugins.huggingface import register_hf_attentions_on_the_fly
        assert register_hf_attentions_on_the_fly in CUSTOM_MODEL_PLUGINS
        CUSTOM_MODEL_PLUGINS.remove(register_hf_attentions_on_the_fly)
        try:
            mtq.quantize(model,cfg,forward_loop=None)
        finally:
            CUSTOM_MODEL_PLUGINS.add(register_hf_attentions_on_the_fly)
        REPORT['quantization']['kv_attention_wrapper_registration']=False
        enabled = [(name,m) for name,m in model.named_modules()
                   if isinstance(m,TensorQuantizer) and m.is_enabled]
        assert len(enabled) == 576, len(enabled)
        assert all(name.rsplit('.',1)[0] in targets for name,m in enabled)
        assert all(getattr(m,'_amax',None) is not None for name,m in enabled if name.endswith('weight_quantizer'))
        REPORT['quantization']['weight_statistics_seconds'] = time.monotonic()-started
        hook.remove()
        REPORT['export_graph_inspection'] = dict(dummy_forward_calls=0,input_tokens=0,quantizers_disabled=True)
        def inspect_probe(module,args):
            assert not any(isinstance(m,TensorQuantizer) and (m.is_enabled or m._if_calib) for m in module.modules())
            assert args[0].shape == (1,2) and torch.equal(args[0],torch.ones_like(args[0]))
            REPORT['export_graph_inspection']['dummy_forward_calls'] += 1
            REPORT['export_graph_inspection']['input_tokens'] += args[0].numel()
        hook = model.register_forward_pre_hook(inspect_probe)
        export_hf_checkpoint(model,dtype=torch.bfloat16,export_dir=OUTPUT,max_shard_size='5GB')
        hook.remove()
    del model,enabled
    gc.collect(); torch.cuda.empty_cache()
    # The upstream Transformers class ignores MTP: retain every original MTP
    # tensor separately, without advertising untested speculative decoding.
    dst_index = json.loads((OUTPUT/'model.safetensors.index.json').read_text())
    mtp = {}
    for name,filename in src_map.items():
        if name.startswith('mtp.'):
            with safe_open(source/filename,framework='pt',device='cpu') as f:
                mtp[name] = f.get_tensor(name).clone()
    assert mtp and not any(name in dst_index['weight_map'] for name in mtp)
    save_file(mtp,str(OUTPUT/'model-mtp.safetensors'),metadata={'format':'pt'})
    dst_index['weight_map'].update({name:'model-mtp.safetensors' for name in mtp})
    dst_index['metadata']['total_size'] += sum(t.numel()*t.element_size() for t in mtp.values())
    (OUTPUT/'model.safetensors.index.json').write_text(json.dumps(dst_index,indent=2)+'\n')
    for file in source.rglob('*'):
        rel=file.relative_to(source)
        if file.is_file() and '.cache' not in rel.parts and not file.name.endswith('.safetensors') and file.name not in ['config.json','model.safetensors.index.json','README.md','.gitattributes']:
            (OUTPUT/rel).parent.mkdir(parents=True,exist_ok=True)
            shutil.copy2(file,OUTPUT/rel)
    # The upstream loader otherwise skips renaming when the branch width is 0.
    # This checkpoint has already folded its branches before quantization.
    for file in OUTPUT.glob('sglang_patch/*/sglang/srt/models/qwen3_5.py'):
        contents=file.read_text()
        before='    if width <= 0:\n        yield from weights\n        return\n'
        after='    if width <= 0:\n        for name, weight in weights:\n            yield name.replace(".delta_attn.", ".linear_attn.").replace(".global_attn.", ".self_attn."), weight\n        return\n'
        assert contents.count(before)==1
        file.write_text(contents.replace(before,after))
    # Both Transformers names and translated serving names must be excluded.
    for filename in ['config.json','hf_quant_config.json']:
        file=OUTPUT/filename
        obj=json.loads(file.read_text())
        quant=obj['quantization_config'] if filename=='config.json' else obj['quantization']
        field = 'ignore' if filename == 'config.json' else 'exclude_modules'
        excludes=list(quant.get(field,[]))
        excludes += ['lm_head','model.visual*','mtp*']
        for i in range(72):
            excludes += [f'model.language_model.layers.{i}.delta_attn*',
                         f'model.language_model.layers.{i}.linear_attn*']
        quant[field]=sorted(set(excludes))
        file.write_text(json.dumps(obj,indent=2)+'\n')
    preserved=0; preserved_count=0
    for filename in sorted(set(src_map.values())):
        with safe_open(source/filename,framework='pt',device='cpu') as sf:
            for name in sf.keys():
                if selected(name) or '.mlp.parallel_ffn.' in name: continue
                assert name in dst_index['weight_map'],name
                with safe_open(OUTPUT/dst_index['weight_map'][name],framework='pt',device='cpu') as df:
                    a,b=sf.get_tensor(name),df.get_tensor(name)
                    assert a.dtype==b.dtype and torch.equal(a,b),name
                    preserved += a.numel()*a.element_size(); preserved_count+=1
    counts=collections.Counter(); packed=0; inputs=0
    for filename in sorted(set(dst_index['weight_map'].values())):
        with safe_open(OUTPUT/filename,framework='pt',device='cpu') as f:
            for name in f.keys():
                tensor=f.get_tensor(name)
                counts[str(tensor.dtype)]+=tensor.numel()*tensor.element_size()
                if selected(name):
                    assert tensor.dtype==torch.uint8,name
                    packed+=1
                if 'scale' in name:
                    assert torch.isfinite(tensor.float()).all() and (tensor.float()>0).all(),name
                if name.endswith('input_scale'):
                    assert torch.equal(tensor,torch.ones_like(tensor)),name
                    inputs+=1
    assert packed==len(targets)==288 and inputs==288,(packed,inputs)
    assert REPORT['quantization']['calibration_forward_calls']==0
    REPORT['export']=dict(source_tensor_bytes=src_index['metadata']['total_size'],
        exported_weight_bytes=sum((OUTPUT/f).stat().st_size for f in set(dst_index['weight_map'].values())),
        tensor_bytes_by_dtype=dict(counts),packed_linear_weights=packed,
        bf16_preserved_bytes=preserved,bf16_preserved_tensors=preserved_count,
        all_unquantized_tensors_bitwise_equal=True,mtp_tensors_preserved=len(mtp),
        activation_global_scales_exactly_one=inputs)
    REPORT['status']='packed_export_verified_pending_native_evaluation'
    record()
    shutil.copy2(ROOT/'quality_report.json',OUTPUT/'quality_report.json')
    API.upload_folder(repo_id=TARGET,folder_path=OUTPUT,
        commit_message='Upload verified plain NVFP4 Agnes Preview checkpoint')
    print('AGNES_BUILD_COMPLETE '+json.dumps(REPORT['export']),flush=True)

if __name__=='__main__':
    try: main()
    except Exception:
        REPORT['status']='build_failed'
        record()
        raise