File size: 2,376 Bytes
12f320c | 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 | import argparse
import json
import os
import time
from pathlib import Path
from reference_eval import REVISION, checkpoint_info, load_reference, sha256, write_json
p = argparse.ArgumentParser()
p.add_argument('--model', type=Path, required=True)
p.add_argument('--input', type=Path, required=True)
p.add_argument('--output', type=Path, required=True)
p.add_argument('--threads', type=int, default=8)
p.add_argument('--max-new-tokens', type=int, default=64)
a = p.parse_args()
os.environ['HF_HUB_OFFLINE'] = '1'
os.environ['TRANSFORMERS_OFFLINE'] = '1'
import torch
from transformers import AutoTokenizer
torch.set_num_threads(a.threads)
torch.set_num_interop_threads(1)
torch.manual_seed(9472)
torch.use_deterministic_algorithms(True)
a.output.mkdir(parents=True, exist_ok=False)
config, tensors, provenance = checkpoint_info(a.model, REVISION)
metadata = {'status': 'loading', 'kind': 'reference_greedy_generation', 'input_sha256': sha256(a.input), 'max_new_tokens': a.max_new_tokens, 'do_sample': False, 'seed': 9472, 'provenance': provenance}
write_json(a.output / 'metadata.json', metadata)
try:
model, loading = load_reference(a.model, tensors, torch)
tokenizer = AutoTokenizer.from_pretrained(a.model, local_files_only=True)
with (a.output / 'generations.jsonl').open('x') as out, torch.inference_mode():
for row in map(json.loads, a.input.read_text().splitlines()):
ids = torch.tensor([row['token_ids']], dtype=torch.long)
start = time.monotonic()
output = model.generate(input_ids=ids, attention_mask=torch.ones_like(ids), do_sample=False, max_new_tokens=a.max_new_tokens, use_cache=True, pad_token_id=tokenizer.eos_token_id)
generated = output[0, ids.shape[1]:].tolist()
record = {'id': row['id'], 'prompt_token_ids': row['token_ids'], 'generated_token_ids': generated, 'text': tokenizer.decode(generated, skip_special_tokens=False), 'seconds': time.monotonic() - start}
out.write(json.dumps(record, ensure_ascii=False) + '\n')
out.flush()
print(json.dumps({'id': row['id'], 'generated_tokens': len(generated), 'seconds': record['seconds']}), flush=True)
metadata['status'] = 'complete'
except BaseException as error:
metadata.update(status='failed', error=str(error))
raise
finally:
write_json(a.output / 'metadata.json', metadata)
|