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)
|