mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
* 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
243 lines
7.5 KiB
Python
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()
|