Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
54 lines
1.5 KiB
Python
54 lines
1.5 KiB
Python
import json
|
|
import sys
|
|
|
|
import torch
|
|
|
|
from kda.training.data import (
|
|
chunk_ids,
|
|
fetch_wiki_texts,
|
|
interleave_balanced,
|
|
split_heldout,
|
|
)
|
|
|
|
|
|
def test_interleave_is_one_to_one_monolingual_blocks():
|
|
a = list(range(10))
|
|
b = list(range(100, 112))
|
|
out = interleave_balanced(a, b, block=4)
|
|
assert out == [0, 1, 2, 3, 100, 101, 102, 103, 4, 5, 6, 7, 104, 105, 106, 107]
|
|
|
|
|
|
def test_split_heldout_keeps_at_least_one_train():
|
|
chunks = torch.arange(10).view(10, 1, 1)
|
|
train, held = split_heldout(chunks, frac=0.01, min_heldout=1)
|
|
assert train.size(0) == 9
|
|
assert held.size(0) == 1
|
|
empty_train, empty_held = split_heldout(chunks[:1], frac=0.5)
|
|
assert empty_train.size(0) == 1
|
|
assert empty_held.size(0) == 0
|
|
|
|
|
|
def test_chunk_ids_drops_tail():
|
|
ids = list(range(10))
|
|
chunks = chunk_ids(ids, batch=2, seq_len=4)
|
|
assert chunks.shape == (1, 2, 4)
|
|
|
|
|
|
def test_wiki_cache_roundtrip(tmp_path, monkeypatch):
|
|
cache = tmp_path / "pretrain"
|
|
cache.mkdir()
|
|
path = cache / "wiki-zh-n2-limit3.jsonl"
|
|
path.write_text(
|
|
"\n".join(json.dumps({"text": f"article {i}"}) for i in range(3)) + "\n",
|
|
encoding="utf-8",
|
|
)
|
|
|
|
def _boom(*_a, **_k):
|
|
raise AssertionError("must not hit the network")
|
|
|
|
fake = type(sys)("datasets")
|
|
fake.load_dataset = _boom
|
|
monkeypatch.setitem(sys.modules, "datasets", fake)
|
|
texts = fetch_wiki_texts(3, lang="zh", cache_dir=cache)
|
|
assert texts == ["article 0", "article 1", "article 2"]
|