Files
K3/tests/integration/test_pretrain_data.py
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

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"]