doc-to-lora/data/store_top_logits.py

126 lines
4.8 KiB
Python

import argparse
import io
import zipfile
from pathlib import Path
from tqdm import tqdm
import torch
from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForCausalLM, PreTrainedTokenizer
# e.g., python data/store_top_logits.py --data "temp/test_data.parquet" --model "Qwen/Qwen3-0.6B"
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--data', required=True,
help='path to a parquet shard of the dataset')
ap.add_argument('--model', required=True,
help='huggingface model name or checkpoint')
ap.add_argument('--max_logits_to_store', type=int, default=100,
help='max number of logits kept per completion')
ap.add_argument('--limit_samples', type=int, default=None,
help='process only the first n dataset rows')
ap.add_argument('--logit_precision', default='bfloat16',
choices=['float32', 'bfloat16', 'float16'],
help='dtype for saved logits')
ap.add_argument('--index_precision', default='int32',
choices=['int64', 'int32', 'int16'],
help='dtype for saved indices')
args = ap.parse_args()
float_map = {
'float32': torch.float32,
'bfloat16': torch.bfloat16,
'float16': torch.float16,
}
int_map = {
'int64': torch.int64,
'int32': torch.int32,
'int16': torch.int16,
}
l_dtype = float_map[args.logit_precision]
i_dtype = int_map[args.index_precision]
ds = load_dataset('parquet', data_files=args.data)['train']
if args.limit_samples and args.limit_samples > 0:
ds = ds.select(range(min(args.limit_samples, len(ds))))
tok: PreTrainedTokenizer = AutoTokenizer.from_pretrained(
args.model, use_fast=True)
model = AutoModelForCausalLM.from_pretrained(args.model)
model.eval()
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.to(device)
base = Path(args.data).with_suffix('')
model_dir = args.model.replace('/', '__')
out_dir = base / model_dir
out_dir.mkdir(parents=True, exist_ok=True)
if args.limit_samples:
zip_name = f'limited_logits_{args.limit_samples}.zip'
else:
zip_name = 'all_logits.zip'
zip_path = out_dir / zip_name
samp_id = 0
with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zf:
for row in tqdm(ds):
ctx = row['context']
qs = row['prompts_level_0']
ans = row['responses_level_0']
q_list = []
a_list = []
vals_list = []
sel_pos_list = []
for q, a in zip(qs, ans):
messages = [
{'role': 'user', 'content': ctx + '\n' + q},
{'role': 'assistant', 'content': a},
]
tpl = tok.apply_chat_template(
messages,
tokenize=True,
add_special_token=False,
truncation=False,
add_generation_prompt=False,
return_dict=True,
)
tokens_to_mask = tok.apply_chat_template(
messages[:-1],
tokenize=True,
add_special_token=False,
truncation=False,
add_generation_prompt=True,
return_dict=True,
)
number_of_tokens_to_mask = len(tokens_to_mask['input_ids'])
assert tokens_to_mask['input_ids'] == tpl['input_ids'][:number_of_tokens_to_mask]
inp = torch.tensor([tpl['input_ids']], device=device)
with torch.no_grad():
logits = model(inp).logits
logits_for_answer = logits[0, number_of_tokens_to_mask - 1:, :]
assert logits_for_answer.shape[0] > 0
assert logits_for_answer.shape[0] + number_of_tokens_to_mask -1 == inp.size(1)
# just in case, logits to store is super large
k = min(args.max_logits_to_store, logits_for_answer.shape[-1])
vals, idxs = torch.topk(logits_for_answer, k, dim=-1)
q_list.append(q)
a_list.append(a)
vals_list.append(vals.cpu())
sel_pos_list.append(idxs.cpu())
record = {
'context': ctx,
'question': q_list,
'answer': a_list,
'top_logits': vals_list,
'top_positions_positions': sel_pos_list,
}
buf = io.BytesIO()
torch.save(record, buf)
samp_id += 1
zf.writestr(f'{samp_id:07d}.pt', buf.getvalue())
if __name__ == '__main__':
main()