| |
| |
| 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) |
| 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) |
| w(f, blk.n2.weight); w(f, blk.n2.bias) |
| w(f, blk.up.weight); w(f, blk.down.weight) |
| w(f, blk.decay) |
| w(f, m.norm.weight); w(f, m.norm.bias) |
| w(f, m.ra.weight); w(f, m.rb.weight) |
| w(f, m.ca); w(f, m.cb) |
| w(f, m.expand.weight); w(f, m.contract.weight) |
| w(f, m.codebook) |
| print('CTRL_BYTES', os.path.getsize(OUT), 'V D L FF A B ncod', V, D, L, FF, A, Bp, ncod, flush=True) |
|
|
| |
| if not os.path.exists(CODONS): |
| m.codons.cpu().numpy().astype(np.uint8).tofile(CODONS) |
| print('CODONS_BYTES', os.path.getsize(CODONS), flush=True) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|