diff --git a/data/create_test_data.py b/data/create_test_data.py new file mode 100644 index 0000000..3559db0 --- /dev/null +++ b/data/create_test_data.py @@ -0,0 +1,51 @@ +import argparse +import pandas as pd + +def make_test_samples(): + return [ + { + "context": "Alice was beginning to get very tired of sitting by her sister on the bank.", + "prompts_level_0": [ + "Who is sitting by Alice?", + "How is Alice feeling at the start?" + ], + "responses_level_0": [ + "Her sister is sitting by her.", + "She is very tired." + ], + }, + { + "context": "debug infoThen she saw a White Rabbit with pink eyes run close by her.", + "prompts_level_0": [ + "What did Alice see run by?", + "What color were its eyes?" + ], + "responses_level_0": [ + "She saw a White Rabbit.", + "Its eyes were pink." + ], + }, + { + "context": "x" * 200, + "prompts_level_0": ["How long is this context?"], + "responses_level_0": ["It is two hundred characters long."], + }, + ] + +def main(output_path: str): + samples = make_test_samples() + df = pd.DataFrame(samples) + df.to_parquet(output_path, index=False) + print(f"Wrote {len(df)} test rows to {output_path}") + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "--output", + "-o", + type=str, + default="test_shard.parquet", + help="Where to write the test Parquet file" + ) + args = parser.parse_args() + main(args.output) diff --git a/data/store_top_logits.py b/data/store_top_logits.py new file mode 100644 index 0000000..bcf2ea3 --- /dev/null +++ b/data/store_top_logits.py @@ -0,0 +1,126 @@ +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()