doc-to-lora/data/generate_ctx_kv.py
Rujikorn Charakorn c5e9bc769d
toy ctx nums and multi-lora training (#11)
* multi-lora trainable toy number repeat dataset

* per rank bias init

* remove head_bias +simplify merge + skip perplexities metric

* ctx_numbers train example

* self-gen ctx numbers example
2025-08-18 18:38:07 +09:00

243 lines
7.5 KiB
Python

import argparse
import json
import os
import random
import uuid
from collections.abc import Iterable
import wonderwords
# -----------------------------
# Simple, editable config knobs
# -----------------------------
# You can change these defaults or pass CLI flags: --key-type / --value-type
DEFAULT_KEY_TYPE = "uuid" # choices: "uuid" | "digits" | "words"
DEFAULT_VALUE_TYPE = "uuid" # choices: "uuid" | "digits" | "words"
PAIR_SEP = "\n" # between pairs
KV_SEP = " : " # between key and value
TOKENS_PER_PAIR = 10 # rough heuristic used for bucketing by context length
BASE_SAMPLES_PER_BIN = 128_000
RNG_SEED = 42
nouns = wonderwords.random_word._get_words_from_text_file("nounlist.txt")
adjs = wonderwords.random_word._get_words_from_text_file("adjectivelist.txt")
# verbs = wonderwords.random_word._get_words_from_text_file("verblist.txt")
words = [f"{adj}-{noun}" for adj in adjs for noun in nouns]
words = sorted(list(set(words)))
def save_jsonl(data: list[dict], filepath: str) -> None:
"""Save data to a JSONL file."""
parent_dir = os.path.dirname(filepath)
if parent_dir:
os.makedirs(parent_dir, exist_ok=True)
with open(filepath, "w") as f:
for entry in data:
json.dump(entry, f)
f.write("\n")
def _gen_uuid8() -> str:
return uuid.uuid4().hex[:8]
def _gen_digits6() -> str:
return f"{random.randint(0, 999_999):06d}"
def _make_generator(kind: str):
kind = kind.lower()
if kind == "uuid":
return _gen_uuid8
if kind == "digits":
return _gen_digits6
if kind == "words":
return _gen_word
raise ValueError(f"Unknown kind '{kind}', expected 'uuid', 'digits', or 'words'")
def _gen_digits4() -> str:
return f"{random.randint(0, 9_999):04d}"
def _gen_word() -> str:
return random.choice(words)
def _make_value_generator(kind: str):
"""Like _make_generator, but when kind is 'digits' use 4 digits for values."""
kind = kind.lower()
if kind == "uuid":
return _gen_uuid8
if kind == "digits":
return _gen_digits4
if kind == "words":
return _gen_word
raise ValueError(f"Unknown kind '{kind}', expected 'uuid', 'digits', or 'words'")
def _unique(seq: Iterable[str]) -> list[str]:
s = set()
out = []
for x in seq:
if x not in s:
s.add(x)
out.append(x)
return out
def render_context(pairs: list[tuple[str, str]]) -> str:
"""Render key-value pairs as a single context string.
Example: "k1:v1 | k2:v2 | k3:v3"
"""
return PAIR_SEP.join([f"{k}{KV_SEP}{v}" for k, v in pairs])
def generate_kv_dataset(n: int, k: int, key_type: str, value_type: str):
"""Generate n examples. Each example has k key-value pairs in the context.
- context: "k1:v1 | k2:v2 | ..."
- prompt: "What is the value of `{k}`?"
- response: "v" (exactly the value corresponding to the chosen k)
"""
key_gen = _make_generator(key_type)
val_gen = _make_value_generator(value_type)
dataset = []
for _ in range(n):
# Generate unique keys to avoid ambiguity
keys: list[str] = []
while len(keys) < k:
keys = _unique([key_gen() for _ in range(k * 2)]) # oversample, then dedupe
keys = keys[:k]
values = [val_gen() for _ in range(k)]
pairs = list(zip(keys, values))
# Choose a random key from the context to ask about
q_key = random.choice(keys)
q_val = values[keys.index(q_key)]
entry = {
"context": f"key{KV_SEP}value\n" + render_context(pairs),
"prompt": f'What is the value of key "{q_key}"? Reply with only the value.',
"response": q_val,
}
dataset.append(entry)
# Split into train/val/test matching the same structure as generate_ctx_numbers.py
total_size = len(dataset)
train_size = int(0.98 * total_size)
val_size = int(0.01 * total_size)
train_data = dataset[:train_size]
val_data = dataset[train_size : train_size + val_size]
test_data = dataset[train_size + val_size :]
return train_data, val_data, test_data
def main():
parser = argparse.ArgumentParser(
description="Generate key-value context QA data with the same structure as generate_ctx_numbers.py",
)
parser.add_argument(
"--key-type",
choices=["uuid", "digits", "words"],
default=DEFAULT_KEY_TYPE,
help="Type of keys to generate: 8-char uuid, 6-digit numbers, or words",
)
parser.add_argument(
"--value-type",
choices=["uuid", "digits", "words"],
default=DEFAULT_VALUE_TYPE,
help="Type of values to generate: 8-char uuid, 4-digit numbers, or words",
)
parser.add_argument("--seed", type=int, default=RNG_SEED, help="Random seed")
parser.add_argument(
"--base-samples-per-bin",
type=int,
default=BASE_SAMPLES_PER_BIN,
help="Baseline number of samples per token bin (scaled by bin width)",
)
parser.add_argument(
"--out-prefix",
type=str,
default="data/raw_datasets/ctx_kv",
help="Output directory prefix (bin range will be appended)",
)
parser.add_argument(
"--tokens-per-pair",
type=int,
default=TOKENS_PER_PAIR,
help="Heuristic tokens per key:value pair for bucketing",
)
parser.add_argument(
"--only-first-n-bins",
type=int,
default=None,
help="For quick tests: only generate the first N token bins",
)
parser.add_argument(
"--dry-run",
action="store_true",
help="Print a small sample and exit without writing files",
)
args = parser.parse_args()
random.seed(args.seed)
# Same token bins as generate_ctx_numbers.py
tok_bins = [(64, 128), (128, 256), (256, 512)] + [
(512 + 256 * i, 512 + 256 * (i + 1)) for i in range(14)
]
if args.only_first_n_bins is not None:
tok_bins = tok_bins[: args.only_first_n_bins]
# Map token bins to pair-length bins using a simple heuristic
len_bins = [
(lo // args.tokens_per_pair, hi // args.tokens_per_pair)
for (lo, hi) in tok_bins
]
# Optional dry-run: build a tiny sample to preview format and exit
if args.dry_run:
train, val, test = generate_kv_dataset(
n=3,
k=max(1, len_bins[0][0] or 1),
key_type=args.key_type,
value_type=args.value_type,
)
print("Sample entry:")
print(json.dumps(train[0], indent=2))
return
for len_bin, tok_bin in zip(len_bins, tok_bins):
bin_size = max(1, len_bin[1] - len_bin[0]) # avoid div-by-zero for tiny bins
save_dir = f"{args.out_prefix}_{tok_bin[0]}_{tok_bin[1]}"
train_data, val_data, test_data = [], [], []
# Iterate over different context lengths (number of pairs)
for k in range(max(1, len_bin[0]), max(1, len_bin[1])):
# scale like the numbers script
per_k = max(1, args.base_samples_per_bin // bin_size)
train, val, test = generate_kv_dataset(
n=per_k, k=k, key_type=args.key_type, value_type=args.value_type
)
train_data += train
val_data += val
test_data += test
os.makedirs(save_dir, exist_ok=True)
save_jsonl(train_data, f"{save_dir}/train.jsonl")
save_jsonl(val_data, f"{save_dir}/val.jsonl")
save_jsonl(test_data, f"{save_dir}/test.jsonl")
print(f"Dataset generated and saved at {save_dir}")
if __name__ == "__main__":
main()