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)