dna-diskchat-2b-v1 / scripts /export_c_dna.py
jaivial's picture
Upload scripts/export_c_dna.py with huggingface_hub
05b9b8d verified
Raw
History Blame Contribute Delete
2.47 kB
# export_c_dna.py - dump DNA-DiskChat-2B controller weights (fp32) + a prompt for
# the C filter runtime. The codon table is already on disk as codons.u8 (128 B/row).
import os, sys, struct, numpy as np, torch
os.environ['CUDA_VISIBLE_DEVICES'] = ''
sys.path.insert(0, '/root/dna')
from model_dna import DnaChat
from infer_dna import load_tok
CK = os.environ.get('DNA_CK', '/root/dna/ckpt/base.pt')
OUT = '/root/dna/ctrl_dna.bin'
CODONS = '/root/dna/codons.u8'
ck = torch.load(CK, map_location='cpu', weights_only=False)
m = DnaChat(**ck['config'])
m.load_state_dict({k.replace('_orig_mod.', ''): v for k, v in ck['model'].items()}, strict=True)
m.eval()
V, D, L, FF, A, Bp = m.vocab, m.d, m.layers, m.ff, m.a, m.b
ncod = m.ncod
def w(f, t): f.write(np.ascontiguousarray(t.detach().float().numpy().ravel(), dtype='<f4').tobytes())
with open(OUT, 'wb') as f:
f.write(b'DNA1'); f.write(struct.pack('<6i', V, D, L, FF, A, Bp))
w(f, m.embed.weight) # [V,D]
for blk in m.blocks:
w(f, blk.n1.weight); w(f, blk.n1.bias)
w(f, blk.proj.weight); w(f, blk.o.weight) # [4D,D], [D,D]
w(f, blk.n2.weight); w(f, blk.n2.bias)
w(f, blk.up.weight); w(f, blk.down.weight) # [2FF,D], [D,FF]
w(f, blk.decay) # [D]
w(f, m.norm.weight); w(f, m.norm.bias)
w(f, m.ra.weight); w(f, m.rb.weight) # [A,D], [B,D]
w(f, m.ca); w(f, m.cb) # [A,D], [B,D]
w(f, m.expand.weight); w(f, m.contract.weight) # [2D,D], [D,2D]
w(f, m.codebook) # [4,256,4]
print('CTRL_BYTES', os.path.getsize(OUT), 'V D L FF A B ncod', V, D, L, FF, A, Bp, ncod, flush=True)
# export codon table if missing
if not os.path.exists(CODONS):
m.codons.cpu().numpy().astype(np.uint8).tofile(CODONS)
print('CODONS_BYTES', os.path.getsize(CODONS), flush=True)
# prompt tokens
tok = load_tok()
prompt = sys.argv[1] if len(sys.argv) > 1 else 'The capital of France is'
ids = [tok.token_to_id('<s>')] + tok.encode(prompt).ids
np.asarray(ids, np.uint16).tofile('/root/dna/prompt.u16')
print('PROMPT_IDS', len(ids), prompt, flush=True)
# PyTorch reference: greedy next-token ids for parity check with the C runtime
from infer_dna import load_model, generate
mm, _ = load_model('cpu')
ref = generate(mm, tok, prompt, n=32, temp=0.0, device='cpu')
print('PYTORCH_REF:', ref, flush=True)