File size: 2,652 Bytes
17f3094
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Generate official BF16 and FP16 PyTorch reference logits with cached decoding."""
import json
from pathlib import Path
import numpy as np
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

torch.set_num_threads(8)
torch.manual_seed(0)
source = Path('models/source')
out = Path('validation'); out.mkdir(exist_ok=True)
tokenizer = AutoTokenizer.from_pretrained(source)
prompts = [
    'Reply with exactly the word READY.',
    'What is 17 * 23? Give only the number.',
    '用一句话解释什么是太阳能。',
    'Write a Python function that returns the sum of a list of numbers.',
    'Return JSON with keys name and count, values Ada and 3. No explanation.',
    'Explain in one sentence why the sky is blue.',
]
fixtures = []
for i, text in enumerate(prompts):
    ids = tokenizer.apply_chat_template([{'role': 'user', 'content': text}],
        add_generation_prompt=True, enable_thinking=False, return_dict=False)
    fixtures.append({'id': str(i), 'prompt': text, 'input_ids': ids})
results = {}
for label, dtype in [('bf16', torch.bfloat16), ('fp16', torch.float16)]:
    print('Loading', label, flush=True)
    model = AutoModelForCausalLM.from_pretrained(source, dtype=dtype,
        attn_implementation='sdpa').to('cuda').eval()
    logits_saved = {}
    for f in fixtures:
        ids = torch.tensor([f['input_ids']], device='cuda')
        cache = None
        steps = []
        continuation = []
        with torch.inference_mode():
            for step in range(9):
                output = model(input_ids=ids, past_key_values=cache, use_cache=True, logits_to_keep=1)
                cache = output.past_key_values
                logits = output.logits[0, -1].float().cpu().numpy()
                steps.append(logits)
                # Both precisions follow the BF16 trajectory for comparable logits.
                next_id = int(np.argmax(logits)) if label == 'bf16' else f['continuation'][step]
                continuation.append(next_id)
                ids = torch.tensor([[next_id]], device='cuda')
                if next_id in (1, 130073):
                    break
        logits_saved[f['id']] = np.stack(steps)
        if label == 'bf16':
            f['continuation'] = continuation
            f['decoded'] = tokenizer.decode(continuation, skip_special_tokens=False)
        print(label, f['id'], f['decoded'], flush=True)
    np.savez(out / f'reference-{label}.npz', **logits_saved)
    del model, output, cache
    torch.cuda.empty_cache()
(out / 'fixtures.json').write_text(json.dumps(fixtures, indent=2, ensure_ascii=False))
print('Saved reference fixtures and logits', flush=True)