| |
| |
| import os, sys, json, torch, torch.nn.functional as F |
| os.environ.setdefault('DNA_CK', '/root/dna/ckpt-sft/sft.pt') |
| sys.path.insert(0, '/root/dna') |
| from infer_dna import load_model, load_tok, block_step, route_token, make_ram_reader |
|
|
| dev = 'cuda' if torch.cuda.is_available() else 'cpu' |
| tok = load_tok() |
| m, cfg = load_model(dev) |
| read_rows = make_ram_reader(m) |
| S = tok.token_to_id('<s>'); E = tok.token_to_id('</s>') |
| U = tok.token_to_id('<|user|>'); A = tok.token_to_id('<|assistant|>') |
| print('CONFIG', json.dumps(cfg), flush=True) |
|
|
| @torch.no_grad() |
| def chat(q, n=80, temp=0.7, rep=1.3): |
| ids = [S, U] + tok.encode('\n' + q).ids + [E, A] + tok.encode('\n').ids |
| s = [torch.zeros(1, m.d, device=dev) for _ in range(m.layers)] |
| feat = None |
| for tid in ids: |
| x = m.embed(torch.tensor([tid], device=dev)) |
| for i, blk in enumerate(m.blocks): x, s[i] = block_step(blk, x, s[i]) |
| feat = m.norm(x + route_token(m, x, read_rows)) |
| out, seen = [], {} |
| for _ in range(n): |
| lo = F.linear(feat, m.embed.weight)[0].float() |
| for t in seen: lo[t] /= rep |
| if temp > 0: |
| nxt = int(torch.multinomial(F.softmax(lo / temp, -1), 1)) |
| else: |
| nxt = int(lo.argmax()) |
| if nxt == E: break |
| out.append(nxt); seen[nxt] = 1 |
| x = m.embed(torch.tensor([nxt], device=dev)) |
| for i, blk in enumerate(m.blocks): x, s[i] = block_step(blk, x, s[i]) |
| feat = m.norm(x + route_token(m, x, read_rows)) |
| return tok.decode(out).strip() |
|
|
| qs = ["What is the capital of France?", |
| "Write one sentence about dogs.", |
| "What is 2+2?", |
| "Give me one tip for learning to code.", |
| "Explain what water is in one sentence.", |
| "List three primary colors.", |
| "Say hello.", |
| "What is the sun?"] |
| samples = [] |
| for q in qs: |
| a = chat(q, temp=0.0) |
| print(f'Q: {q}\nA: {a}\n---', flush=True) |
| samples.append({'q': q, 'a': a}) |
| json.dump(samples, open('/root/dna/sft_quality.json', 'w'), indent=2) |
| print('SFT_QUALITY_DONE', flush=True) |
|
|