mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
Store top k logits tests and scripts (#2) (to be refactored + integrated with data pipeline)
This commit is contained in:
parent
f5e2f55ada
commit
64f7b3d401
2 changed files with 177 additions and 0 deletions
51
data/create_test_data.py
Normal file
51
data/create_test_data.py
Normal file
|
|
@ -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": "<think>debug info</think>Then 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)
|
||||
126
data/store_top_logits.py
Normal file
126
data/store_top_logits.py
Normal file
|
|
@ -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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue