mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
LoRA separate bias param + merger + toy ds (#15)
* separate bias + head weight init * smaller bias init + sum with 1/sqrt(combined_r) scaling + better lora merge logic * ok init * chunked ctx magic num working * configs * no transpose + only sum aggregation * fix cli * cli * 32-256 data * example cli + smaller dataset * toy ctx magic num ds size + quantize + remove head bias * configs removed * require lora_bias * Merge remote-tracking branch 'origin/main' into new-arch
This commit is contained in:
parent
9d82014dc2
commit
ee9039295e
18 changed files with 400 additions and 455 deletions
|
|
@ -8,7 +8,9 @@ import random
|
|||
# Config knobs (edit or use CLI)
|
||||
# -----------------------------
|
||||
TOKENS_PER_BLOCK = 40 # rough heuristic tokens per noise block
|
||||
BASE_SAMPLES_PER_BIN = 12_800
|
||||
BASE_SAMPLES_PER_BIN = (
|
||||
320_000 # training samples budget scaler only (val/test fixed at 1000 each)
|
||||
)
|
||||
RNG_SEED = 42
|
||||
NOISE_BLOCK = "The grass is green. The sky is blue. The sun is yellow. Here we go. There and back again."
|
||||
SPECIAL_TPL = "The special magic number is {magic_number}."
|
||||
|
|
@ -82,29 +84,19 @@ def _build_example(total_blocks: int, depth_bin: int) -> dict:
|
|||
return {"context": context, "prompt": prompt, "response": response}
|
||||
|
||||
|
||||
def generate_magic_dataset(n: int, k: int) -> tuple[list[dict], list[dict], list[dict]]:
|
||||
"""Generate n samples for a given block length k, evenly distributed across 10 depth bins."""
|
||||
# Evenly divide n across 10 depth bins
|
||||
def generate_examples(n: int, k: int) -> list[dict]:
|
||||
"""Generate n examples (all for block length k) evenly across 10 depth bins."""
|
||||
if n <= 0:
|
||||
return []
|
||||
base = n // 10
|
||||
rem = n % 10
|
||||
counts = [base + (1 if i < rem else 0) for i in range(10)]
|
||||
|
||||
dataset: list[dict] = []
|
||||
out: list[dict] = []
|
||||
for depth_bin, c in enumerate(counts):
|
||||
for _ in range(c):
|
||||
dataset.append(_build_example(total_blocks=k, depth_bin=depth_bin))
|
||||
|
||||
# Randomize before splitting to ensure val/test are random samples
|
||||
random.shuffle(dataset)
|
||||
|
||||
# 98/1/1 split
|
||||
total = len(dataset)
|
||||
train_sz = int(0.98 * total)
|
||||
val_sz = int(0.01 * total)
|
||||
train = dataset[:train_sz]
|
||||
val = dataset[train_sz : train_sz + val_sz]
|
||||
test = dataset[train_sz + val_sz :]
|
||||
return train, val, test
|
||||
out.append(_build_example(total_blocks=k, depth_bin=depth_bin))
|
||||
random.shuffle(out)
|
||||
return out
|
||||
|
||||
|
||||
def main():
|
||||
|
|
@ -112,11 +104,17 @@ def main():
|
|||
description="Generate noise-wrapped special magic number dataset (similar structure to generate_ctx_kv.py)",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=RNG_SEED, help="Random seed")
|
||||
parser.add_argument(
|
||||
"--tokenizer-name",
|
||||
type=str,
|
||||
default="google/gemma-2-2b-it",
|
||||
help=("Tokenizer name"),
|
||||
)
|
||||
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)",
|
||||
help="Baseline number of TRAINING samples per token bin (scaled by bin width). Validation & test are always 1000 each.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--out-prefix",
|
||||
|
|
@ -148,47 +146,140 @@ def main():
|
|||
|
||||
random.seed(args.seed)
|
||||
|
||||
# Same token bins pattern as generate_ctx_kv.py (current version)
|
||||
tok_bins = [(64, 128), (128, 256), (256, 512)] + [
|
||||
(512 + 256 * i, 512 + 256 * (i + 1)) for i in range(14)
|
||||
]
|
||||
# ----------------------------------------------------
|
||||
# Optional: report tokenizer-based token length stats
|
||||
# ----------------------------------------------------
|
||||
if args.tokenizer_name:
|
||||
try:
|
||||
from transformers import AutoTokenizer # type: ignore
|
||||
except Exception as e: # pragma: no cover
|
||||
raise RuntimeError(
|
||||
"Failed to import transformers. Install it or omit --tokenizer-name."
|
||||
) from e
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_name, use_fast=True)
|
||||
noise_token_count = len(tokenizer(NOISE_BLOCK).input_ids)
|
||||
special_example = SPECIAL_TPL.format(magic_number="0000")
|
||||
special_token_count = len(tokenizer(special_example).input_ids)
|
||||
print(
|
||||
f"[Tokenizer: {args.tokenizer_name}] Noise block tokens: {noise_token_count} | Special line tokens: {special_token_count}"
|
||||
)
|
||||
|
||||
tok_bins = [(32, 128), (128, 256), (256, 512), (512, 1024), (32, 1024)] + [
|
||||
(1024 * i, 1024 * (i + 1)) for i in range(1, 16)
|
||||
]
|
||||
if args.only_first_n_bins is not None:
|
||||
tok_bins = tok_bins[: args.only_first_n_bins]
|
||||
|
||||
# Map token bins to block-length bins using heuristic
|
||||
len_bins = [
|
||||
(lo // args.tokens_per_block, hi // args.tokens_per_block)
|
||||
for (lo, hi) in tok_bins
|
||||
]
|
||||
if args.tokenizer_name:
|
||||
max_hi = max(hi for _, hi in tok_bins)
|
||||
|
||||
def measure_len(k: int) -> int:
|
||||
if k == 1:
|
||||
ctx = SPECIAL_TPL.format(magic_number="0000")
|
||||
else:
|
||||
blocks = [NOISE_BLOCK] * (k - 1) + [
|
||||
SPECIAL_TPL.format(magic_number="0000")
|
||||
]
|
||||
ctx = SEP.join(blocks)
|
||||
return len(tokenizer(ctx).input_ids)
|
||||
|
||||
lengths: list[int] = [0]
|
||||
k = 1
|
||||
while True:
|
||||
L = measure_len(k)
|
||||
lengths.append(L)
|
||||
if L >= max_hi:
|
||||
break
|
||||
k += 1
|
||||
|
||||
len_bins = []
|
||||
for lo, hi in tok_bins:
|
||||
k_lo = None
|
||||
for kk in range(1, len(lengths)):
|
||||
if lengths[kk] >= lo:
|
||||
k_lo = kk
|
||||
break
|
||||
if k_lo is None or lengths[k_lo] >= hi:
|
||||
len_bins.append((0, 0))
|
||||
continue
|
||||
k_hi = len(lengths)
|
||||
for kk in range(k_lo, len(lengths)):
|
||||
if lengths[kk] >= hi:
|
||||
k_hi = kk
|
||||
break
|
||||
len_bins.append((k_lo, k_hi))
|
||||
|
||||
base_tokens = lengths[1]
|
||||
delta = (lengths[2] - lengths[1]) if len(lengths) > 2 else 0
|
||||
print(
|
||||
f"Using tokenizer-measured block ranges. base_tokens={base_tokens} approx_delta={delta}"
|
||||
)
|
||||
else:
|
||||
len_bins = [
|
||||
(lo // args.tokens_per_block, hi // args.tokens_per_block)
|
||||
for (lo, hi) in tok_bins
|
||||
]
|
||||
|
||||
if args.dry_run:
|
||||
k = max(1, len_bins[0][0] or 1)
|
||||
train, _, _ = generate_magic_dataset(n=10, k=k)
|
||||
print("Sample entry:")
|
||||
print(json.dumps(train[0], indent=2))
|
||||
for lb in len_bins:
|
||||
if lb[1] > lb[0]:
|
||||
k = max(1, lb[0])
|
||||
sample = generate_examples(10, k)
|
||||
print("Sample entry:")
|
||||
print(json.dumps(sample[0], indent=2))
|
||||
break
|
||||
return
|
||||
|
||||
# -----------------------------------------------
|
||||
# Main generation per token bin
|
||||
# -----------------------------------------------
|
||||
TARGET_VAL = 1000
|
||||
TARGET_TEST = 1000
|
||||
for len_bin, tok_bin in zip(len_bins, tok_bins):
|
||||
bin_size = max(1, len_bin[1] - len_bin[0])
|
||||
if len_bin[1] <= len_bin[0]:
|
||||
print(f"Skipping token bin {tok_bin} (no valid block counts)")
|
||||
continue
|
||||
k_start = max(1, len_bin[0])
|
||||
k_end = max(1, len_bin[1])
|
||||
k_values = list(range(k_start, k_end))
|
||||
bin_size = len(k_values)
|
||||
save_dir = f"{args.out_prefix}_{tok_bin[0]}_{tok_bin[1]}"
|
||||
train_data: list[dict] = []
|
||||
training_enabled = tok_bin[1] <= 512 # unchanged policy
|
||||
if training_enabled:
|
||||
train_data: list[dict] = []
|
||||
# Distribute training budget across k values.
|
||||
# Scale: per_k = base_samples_per_bin / bin_size
|
||||
per_k_train = max(1, args.base_samples_per_bin // max(1, bin_size))
|
||||
for k in k_values:
|
||||
train_data += generate_examples(per_k_train, k)
|
||||
val_data: list[dict] = []
|
||||
test_data: list[dict] = []
|
||||
|
||||
for k in range(max(1, len_bin[0]), max(1, len_bin[1])):
|
||||
per_k = max(1, args.base_samples_per_bin // bin_size)
|
||||
train, val, test = generate_magic_dataset(n=per_k, k=k)
|
||||
train_data += train
|
||||
val_data += val
|
||||
test_data += test
|
||||
|
||||
base_val = TARGET_VAL // bin_size
|
||||
rem_val = TARGET_VAL % bin_size
|
||||
base_test = TARGET_TEST // bin_size
|
||||
rem_test = TARGET_TEST % bin_size
|
||||
for idx, k in enumerate(k_values):
|
||||
n_val_k = base_val + (1 if idx < rem_val else 0)
|
||||
n_test_k = base_test + (1 if idx < rem_test else 0)
|
||||
if n_val_k:
|
||||
val_data += generate_examples(n_val_k, k)
|
||||
if n_test_k:
|
||||
test_data += generate_examples(n_test_k, k)
|
||||
random.shuffle(val_data)
|
||||
random.shuffle(test_data)
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
save_jsonl(train_data, f"{save_dir}/train.jsonl")
|
||||
if training_enabled:
|
||||
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 training_enabled:
|
||||
print(
|
||||
f"Dataset generated at {save_dir} (train={len(train_data)} val={len(val_data)} test={len(test_data)})"
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"Dataset (val/test only) generated at {save_dir} (val={len(val_data)} test={len(test_data)})"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue