#!/usr/bin/env python3
"""BF16 -> NVFP4 (compressed-tensors, nvfp4-pack-quantized) for Qwen3.8-27B, CPU only.

Reproduces the layout of lyf/Qwen3.8-27B-Heretic-ARA-NVFP4-MTP-VL (ModelOpt NVFP4_DEFAULT_CFG):
  - every linear layer -> NVFP4: E2M1 4-bit values, one FP8-E4M3 scale per 16 values,
    one FP32 global scale per tensor; q/k/v, the four GDN in_proj and MLP gate/up share a global scale
  - kept in BF16: lm_head, embed_tokens, visual, conv1d, mtp
  - input_global_scale is written as 1.0 (FastLLM on sm_75 does not read it; no calibration)

Usage:
  hf download lyf/Qwen3.8-27B-Heretic-ARA-NVFP4-MTP-VL --include "*.json" "recipe.yaml" --local-dir ./nvfp4-ref
  python3 convert.py /path/to/Qwen3.8-27B-BF16 /path/to/output --reference ./nvfp4-ref

Needs numpy only. Written for ai-muninn.com (2026-09-28); verified bit-exact against the lyf weights on
attention / GDN / MLP layers. No warranty.
"""
import argparse, json, mmap, os, shutil, struct, re
from pathlib import Path
import numpy as np
SIZES={'BF16':2,'F16':2,'F32':4,'F64':8,'I64':8,'I32':4,'U8':1,'F8_E4M3':1}
B4=np.array([.25,.75,1.25,1.75,2.5,3.5,5],np.float32)
V8=np.array([i*2**-9 if i<8 else (1+(i&7)/8)*2**((i>>3)-7) for i in range(127)],np.float32)
M8=(V8[:-1].astype(np.float64)+V8[1:].astype(np.float64))/2
class ST:
 def __init__(self,p):
  self.f=open(p,'rb');self.mm=mmap.mmap(self.f.fileno(),0,access=mmap.ACCESS_READ)
  n=struct.unpack_from('<Q',self.mm)[0];self.base=8+n;self.h=json.loads(self.mm[8:self.base])
 def raw(self,k):
  a,b=self.h[k]['data_offsets'];return memoryview(self.mm)[self.base+a:self.base+b]
 def close(self):self.mm.close();self.f.close()
def f32(raw,shape,dtype):
 if dtype=='BF16':return (np.frombuffer(raw,'<u2').astype('<u4')<<16).view('<f4').reshape(shape)
 if dtype=='F16':return np.frombuffer(raw,'<f2').astype('<f4').reshape(shape)
 if dtype=='F32':return np.frombuffer(raw,'<f4').reshape(shape)
 raise ValueError(dtype)
def e8(x):
 x=np.clip(x,2**-9,448).astype(np.float32)
 i=np.searchsorted(M8,x,side='left').astype(np.uint8)
 i+=( (x==M8[np.minimum(i,len(M8)-1)]) & ((i&1)==1) ).astype(np.uint8)
 return i
def e4(x):
 a=np.abs(x);i=np.searchsorted(B4,a,side='left').astype(np.uint8)
 i+=((a==.75)|(a==1.75)|(a==3.5)).astype(np.uint8)
 return i|(np.signbit(x).astype(np.uint8)<<3)
def quant(raw,shape,dtype,global_g=None):
 assert len(shape)==2 and shape[1]%16==0,(shape,dtype)
 w=f32(raw,shape,dtype);g=np.float32(np.max(np.abs(w))/np.float32(2688)) if global_g is None else np.float32(global_g)
 if g==0:g=np.float32(1)
 packed=bytearray(w.size//2);scales=bytearray(w.size//16)
 cols=shape[1]
 for a in range(0,shape[0],128):
  b=min(shape[0],a+128);x=w[a:b].reshape(b-a,-1,16)
  s=np.max(np.abs(x),axis=-1)/(np.float32(6)*g)
  s=np.where(s==0,np.float32(1),s);sb=e8(s)
  scales[a*cols//16:b*cols//16]=sb.tobytes()
  q=e4(x/(V8[sb]*g)[...,None]);pq=(q[...,::2]|(q[...,1::2]<<4)).astype(np.uint8)
  packed[a*cols//2:b*cols//2]=pq.tobytes()
 return bytes(packed),bytes(scales),struct.pack('<f',np.float32(1/g))
def shared_group(key):
 if key.startswith("mtp.") or "visual" in key:return None
 patterns=(r"(.*\.self_attn)\.(q_proj|k_proj|v_proj)\.weight$",
           r"(.*\.linear_attn)\.(in_proj_a|in_proj_b|in_proj_qkv|in_proj_z)\.weight$",
           r"(.*\.mlp)\.(gate_proj|up_proj)\.weight$")
 for pattern in patterns:
  m=re.fullmatch(pattern,key)
  if m:return (pattern,m.group(1))
 return None

def shared_global_g(group,source_index,source_files):
 pattern,parent=group
 maximum=np.float32(0)
 for key in source_index:
  m=re.fullmatch(pattern,key)
  if m and m.group(1)==parent:
   st=source_files[source_index[key]];meta=st.h[key]
   value=np.max(np.abs(f32(st.raw(key),meta["shape"],meta["dtype"])))
   maximum=np.maximum(maximum,value)
 return np.float32(maximum/np.float32(2688))

def target_desc(key,meta):
 if key.endswith('.weight') and len(meta['shape'])==2 and not any(z in key for z in ['visual','conv1d']) and not key.startswith('mtp.') and key not in ['lm_head.weight','model.language_model.embed_tokens.weight']:
  pre=key[:-7];r,c=meta['shape'];assert c%16==0,(key,c)
  return [(pre+'.weight_packed','U8',[r,c//2]),(pre+'.weight_scale','F8_E4M3',[r,c//16]),(pre+'.weight_global_scale','F32',[]),(pre+'.input_global_scale','F32',[])]
 return [(key,meta['dtype'],meta['shape'])]
def header(entries):
 h={};pos=0
 for k,t,sh in entries:
  n=SIZES[t]*int(np.prod(sh,dtype=np.int64));h[k]={'dtype':t,'shape':sh,'data_offsets':[pos,pos+n]};pos+=n
 b=json.dumps(h,separators=(',',':'),ensure_ascii=False).encode();b+=b' '*(-len(b)%8)
 return struct.pack('<Q',len(b))+b,pos

def run(src,dst,ref):
 src=Path(src);dst=Path(dst);ref=Path(ref);dst.mkdir(parents=True,exist_ok=True)
 idx=json.load(open(src/'model.safetensors.index.json'))['weight_map']
 layout=json.load(open(ref/'model.safetensors.index.json'))['weight_map']
 files={n:ST(src/n) for n in set(idx.values())}
 try:
  jobs={};outmap={};gcache={}
  for key in sorted(idx):
   meta=files[idx[key]].h[key]
   for name,dtype,shape in target_desc(key,meta):jobs[name]=(key,dtype,shape);outmap[name]=layout.get(name)
  assert set(jobs)==set(layout),(len(jobs),len(layout),list(set(jobs)-set(layout))[:5],list(set(layout)-set(jobs))[:5])
  for shard in sorted(set(layout.values())):
   entries=[(k,jobs[k][1],jobs[k][2]) for k in sorted(jobs) if outmap[k]==shard]
   h,total=header(entries);print('writing',shard,total,flush=True)
   with open(dst/shard,'wb') as f:
    f.write(h);cachekey=None;cache=None
    for name,dtype,shape in entries:
     orig=jobs[name][0];st=files[idx[orig]]
     if orig==name:f.write(st.raw(orig));continue
     if orig!=cachekey:
      group=shared_group(orig)
      if group is not None and group not in gcache:gcache[group]=shared_global_g(group,idx,files)
      cache=quant(st.raw(orig),st.h[orig]["shape"],st.h[orig]["dtype"],gcache.get(group));cachekey=orig
     typ=name.rsplit('.',1)[-1]
     f.write(cache[{'weight_packed':0,'weight_scale':1,'weight_global_scale':2}[typ]] if typ!='input_global_scale' else struct.pack('<f',1.0))
    f.flush();os.fsync(f.fileno())
   print('done',shard,(dst/shard).stat().st_size,flush=True)
  (dst/'model.safetensors.index.json').write_text(json.dumps({'metadata':{'total_size':sum(SIZES[dtype]*int(np.prod(shape,dtype=np.int64)) for orig,dtype,shape in jobs.values())},'weight_map':outmap},indent=2)+'\n')
  for p in src.iterdir():
   if p.is_file() and p.name not in ['config.json','model.safetensors.index.json'] and not p.name.endswith('.safetensors'):shutil.copy2(p,dst/p.name)
  cfg=json.load(open(src/'config.json'));cfg['quantization_config']=json.load(open(ref/'config.json'))['quantization_config']
  (dst/'config.json').write_text(json.dumps(cfg,indent=2,ensure_ascii=False)+'\n')
  for n in ['hf_quant_config.json','recipe.yaml']:shutil.copy2(ref/n,dst/n)
 finally:
  for st in files.values():st.close()
if __name__=='__main__':
 p=argparse.ArgumentParser();p.add_argument('source');p.add_argument('dest');p.add_argument('--reference',required=True,help='dir with the lyf NVFP4 repo json/yaml files');a=p.parse_args();run(a.source,a.dest,a.reference)
