File size: 9,080 Bytes
19c1f65
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
#!/usr/bin/env python3
"""Reproduce selected-result prompting versus the frozen STRATA response interface."""

import argparse
from collections import defaultdict
import gzip
import json
import os
from pathlib import Path
import subprocess
import sys
import time

ROOT = Path(__file__).resolve().parents[1]
DATA = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT))

def selected_prompt(record, config):
    selected = {key: record[key] for key in config['selected_fields']}
    return config['prompt_prefix'] + json.dumps(selected, ensure_ascii=False, separators=(',', ':')) + config['prompt_suffix'] + record['query']

def exact_value_once(text, value):
    return text.encode('utf-8').count(value.encode('utf-8')) == 1

def make_infer(torch, model, head, tokenizer, rows, records, config, system,
               state_table, banks, versions, device):
    from strata.eval.native_lm_benchmark import _prompt_ids
    from strata.eval.native_lm_frame_separated_copy import frame_separated_generate
    backbone = model.backbone

    @torch.inference_mode()
    def infer(arm, indices):
        group = [rows[i] for i in indices]
        torch.cuda.synchronize()
        start = time.perf_counter()
        if arm == 'PROMPT_SERIALIZATION':
            sequences = [_prompt_ids(tokenizer, selected_prompt(records[i], config)) for i in indices]
            width = max(map(len, sequences))
            ids = torch.tensor([[tokenizer.pad_token_id] * (width - len(seq)) + seq for seq in sequences], device=device)
            mask = torch.tensor([[0] * (width - len(seq)) + [1] * len(seq) for seq in sequences], device=device)
            outputs = backbone.generate(input_ids=ids, attention_mask=mask, do_sample=False, use_cache=True, max_new_tokens=config['prompt_generation_max_new_tokens'], pad_token_id=tokenizer.pad_token_id, eos_token_id=tokenizer.eos_token_id)
            texts = tokenizer.batch_decode(outputs[:, width:], skip_special_tokens=True)
            generated_lengths = [len(tokens) - list(tokens).count(tokenizer.pad_token_id) for tokens in outputs[:, width:].cpu().tolist()]
            truncated = [tokenizer.eos_token_id not in tokens for tokens in outputs[:, width:].cpu().tolist()]
        else:
            outputs, _ = frame_separated_generate(model, head, tokenizer, group, state_table, [banks[i] for i in indices], [0] * len(group), batch_size=len(group), max_actions=config['strata_max_actions'], frame_handle=system['frame']['canonical_frame_handle'], frame_surrogate=system['frame']['canonical_frame_surrogate'], terminator=system['frame']['structural_terminator'], current_versions=[versions[i] for i in indices])
            texts = [output.text for output in outputs]
            generated_lengths = [sum((action.startswith('GEN(') for action in output.actions)) for output in outputs]
            truncated = [output.status == 'MAX_ACTIONS' for output in outputs]
        torch.cuda.synchronize()
        elapsed = time.perf_counter() - start
        lengths = [len(_prompt_ids(tokenizer, selected_prompt(records[i], config) if arm == 'PROMPT_SERIALIZATION' else records[i]['query'])) for i in indices]
        return (texts, lengths, generated_lengths, truncated, elapsed)
    return infer

def read_records():
    with gzip.open(DATA/'records.jsonl.gz', 'rt', encoding='utf-8') as handle:
        return [json.loads(line) for line in handle]


def score(paths):
    gold = {row['id']:row for row in read_records()}
    scores = defaultdict(lambda: {'correct':0,'total':0,'input_tokens':0,'truncated':0})
    timings = {}
    for path in paths:
        opener = gzip.open if str(path).endswith('.gz') else open
        with opener(path, 'rt', encoding='utf-8') as handle:
            for line in handle:
                item = json.loads(line)
                row = gold[item['id']]
                value = scores[item['arm']]
                value['correct'] += int(exact_value_once(item['text'],row['value']))
                value['total'] += 1
                value['input_tokens'] += item['input_tokens']
                value['truncated'] += item['truncated']
                timings[(item['arm'],item['rank'],item['repeat'],item['batch_offset'])] = item['batch_seconds']
    for arm,value in scores.items():
        value['accuracy'] = value['correct']/value['total']
        value['mean_input_tokens'] = value['input_tokens']/value['total']
        value['amortized_ms_per_record'] = sum(t for key,t in timings.items() if key[0]==arm)*1000/value['total']
    return dict(scores)


def worker(args, config):
    import torch
    from load_and_answer import load_model
    from strata.data.native_lm_integration import NativeLMExample, address_codes, answer_text
    from strata.modeling.exact_payload_realizer import PayloadAuthority
    from strata.training.native_lm_integration import compact_state_table
    torch.set_num_threads(2)
    torch.cuda.set_device(args.rank)
    torch.manual_seed(config['seed'])
    device = torch.device('cuda',args.rank)
    system,tokenizer,model,head,codec = load_model(ROOT,args.base_model,device)
    state_table = compact_state_table(codec)
    records = [row for row in read_records() if row['rank']==args.rank]
    if args.limit is not None:
        records = records[:args.limit]
    rows = [NativeLMExample(example_id=r['id'],split='system-v1',schema=r['schema'],field=r['field'],
                event=r['event'],predicate=r['predicate'],role=r['role'],value_type=r['value_type'],
                payload_handle=r['payload_handle'],value=r['value'],
                address_codes=address_codes(r['event'],r['predicate'],r['role']),query=r['query'],
                full_history_query='',answer=answer_text(r['field'],r['value']),operation=r['operation'],age_windows=0)
            for r in records]
    versions = [r['event_version'] for r in records]
    banks = [[PayloadAuthority.issue(event=r.event,predicate=r.predicate,role=r.role,
                handle=r.payload_handle,payload=r.value,version=v)] for r,v in zip(rows,versions)]
    infer = make_infer(torch,model,head,tokenizer,rows,records,config,system,state_table,banks,versions,device)
    warm = list(range(min(config['batch_size'],len(rows))))
    for arm in ['PROMPT_SERIALIZATION','STRATA_NATIVE_V1']:
        infer(arm,warm)
    path = args.output/f'rank-{args.rank}.jsonl'
    with path.open('x',encoding='utf-8') as handle:
        for repeat in range(config['repeats']):
            for offset in range(0,len(rows),config['batch_size']):
                indices = list(range(offset,min(offset+config['batch_size'],len(rows))))
                arms = ['PROMPT_SERIALIZATION','STRATA_NATIVE_V1']
                if (args.rank+repeat+offset//config['batch_size'])%2:
                    arms.reverse()
                for arm in arms:
                    texts,lengths,generated,truncated,elapsed = infer(arm,indices)
                    for index,text,length,n,cutoff in zip(indices,texts,lengths,generated,truncated):
                        handle.write(json.dumps({'id':records[index]['id'],'rank':args.rank,'repeat':repeat,
                            'arm':arm,'text':text,'input_tokens':length,'generated_tokens':n,'truncated':cutoff,
                            'batch_offset':offset,'batch_seconds':elapsed,'batch_size':len(indices)},ensure_ascii=False)+'\n')


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--score-published',action='store_true')
    parser.add_argument('--base-model',default=os.environ.get('STRATA_BASE_MODEL'))
    parser.add_argument('--output',type=Path)
    parser.add_argument('--rank',type=int,choices=range(8))
    parser.add_argument('--limit',type=int)
    args = parser.parse_args()
    if args.score_published:
        print(json.dumps(score([DATA/'responses.jsonl.gz']),indent=2))
        return
    if not args.base_model or not args.output:
        parser.error('--base-model and --output are required for inference')
    config = json.loads((DATA/'protocol.json').read_text())
    if args.rank is not None:
        args.output.mkdir(parents=True,exist_ok=True)
        worker(args,config)
        return
    args.output.mkdir(parents=True,exist_ok=False)
    processes = []
    for rank in range(8):
        log = (args.output/f'rank-{rank}.log').open('x')
        command = [sys.executable,__file__,'--rank',str(rank),'--base-model',args.base_model,'--output',str(args.output)]
        if args.limit is not None:
            command += ['--limit',str(args.limit)]
        processes.append((subprocess.Popen(command,stdout=log,stderr=subprocess.STDOUT,
            env={**os.environ,'OMP_NUM_THREADS':'2','TOKENIZERS_PARALLELISM':'false'}),log))
    codes = []
    for process,log in processes:
        codes.append(process.wait())
        log.close()
    if any(codes):
        raise SystemExit(f'Inference failed: {codes}')
    results = score(sorted(args.output.glob('rank-*.jsonl')))
    (args.output/'results.json').write_text(json.dumps(results,indent=2))
    print(json.dumps(results,indent=2))


if __name__ == '__main__':
    main()