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)